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,230 @@
1
+ """Checkpoint downloading and pretrained weight loading utilities for MatGL models."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ import json
7
+ import logging
8
+ from pathlib import Path
9
+ from typing import Optional, Union, Dict, Any, List
10
+ import numpy as np
11
+ import torch
12
+ import keras
13
+
14
+ logger = logging.getLogger(__name__)
15
+
16
+ HF_MATGL_ORG = "materialyze"
17
+
18
+ KNOWN_PRETRAINED_MODELS = [
19
+ "CHGNet-PES-MatPES-PBE-2025.2.10",
20
+ "CHGNet-PES-MatPES-r2SCAN-2025.2.10",
21
+ "M3GNet-Eform-MP-2018.6.1",
22
+ "M3GNet-Eform-MP-2019.4.1",
23
+ "M3GNet-PES-ANI-1x-Subset",
24
+ "M3GNet-PES-MatPES-PBE-2025.2",
25
+ "M3GNet-PES-MatPES-r2SCAN-2025.2",
26
+ "MEGNet-BandGap-mfi-MP-2019.4.1",
27
+ "MEGNet-Eform-MP-2018.6.1",
28
+ "QET-PES-MatPES-PBE-2025.2",
29
+ "QET-PES-MatPES-r2SCAN-2025.2",
30
+ "QET-PES-MatQ",
31
+ "SO3Net-PES-ANI-1x-Subset",
32
+ "TensorNet-PES-ANI-1x-Subset",
33
+ "TensorNet-PES-MatPES-PBE-2025.2",
34
+ "TensorNet-PES-MatPES-PBE-2025.2-m",
35
+ "TensorNet-PES-MatPES-r2SCAN-2025.2",
36
+ "TensorNet-PES-MatPES-r2SCAN-2025.2-m",
37
+ ]
38
+
39
+
40
+ def get_available_pretrained_models() -> List[str]:
41
+ """Return list of available pretrained materials models."""
42
+ try:
43
+ from huggingface_hub import HfApi
44
+ api = HfApi()
45
+ names = []
46
+ for model_info in api.list_models(author=HF_MATGL_ORG):
47
+ repo_id = str(getattr(model_info, "id", "") or getattr(model_info, "modelId", "") or "")
48
+ if "/" in repo_id:
49
+ names.append(repo_id.split("/", 1)[1])
50
+ if names:
51
+ return sorted(names)
52
+ except Exception as e:
53
+ logger.debug("Hugging Face API listing unavailable, returning static list: %s", e)
54
+ return sorted(KNOWN_PRETRAINED_MODELS)
55
+
56
+
57
+ def download_matgl_checkpoint(
58
+ name_or_repo: str,
59
+ folder: str = "checkpoints",
60
+ log: bool = True,
61
+ ) -> Dict[str, str]:
62
+ """Download matgl model files (model.json, state.pt) from Hugging Face Hub."""
63
+ from huggingface_hub import hf_hub_download
64
+
65
+ repo_id = name_or_repo if "/" in name_or_repo else f"{HF_MATGL_ORG}/{name_or_repo}"
66
+ if log:
67
+ print(f"Downloading checkpoint files for {repo_id}...")
68
+
69
+ os.makedirs(folder, exist_ok=True)
70
+ f_json = hf_hub_download(repo_id=repo_id, filename="model.json", cache_dir=folder)
71
+ f_state = hf_hub_download(repo_id=repo_id, filename="state.pt", cache_dir=folder)
72
+
73
+ return {
74
+ "model.json": f_json,
75
+ "state.pt": f_state,
76
+ "repo_id": repo_id,
77
+ }
78
+
79
+
80
+ def _assign_dense_weights(keras_layer, weight_tensor, bias_tensor=None):
81
+ """Assign PyTorch Linear weight/bias to Keras Dense layer."""
82
+ w_np = weight_tensor.detach().cpu().numpy()
83
+ if len(w_np.shape) == 2:
84
+ # PyTorch is [out_features, in_features], Keras is [in_features, out_features]
85
+ w_np = np.transpose(w_np, (1, 0))
86
+ weights = [w_np]
87
+ if bias_tensor is not None:
88
+ b_np = bias_tensor.detach().cpu().numpy()
89
+ weights.append(b_np)
90
+ keras_layer.set_weights(weights)
91
+
92
+
93
+ def load_matgl_weights(
94
+ model: keras.Model,
95
+ state_dict_or_path: Union[str, Dict[str, torch.Tensor]],
96
+ log: bool = False,
97
+ ) -> int:
98
+ """Load PyTorch checkpoint weights into multi-backend Keras 3 MatGL model."""
99
+ if isinstance(state_dict_or_path, (str, Path)):
100
+ state = torch.load(state_dict_or_path, map_location="cpu", weights_only=True)
101
+ else:
102
+ state = state_dict_or_path
103
+
104
+ # Strip potential 'model.' prefix
105
+ cleaned_state = {}
106
+ for k, v in state.items():
107
+ if k.startswith("model."):
108
+ cleaned_state[k[6:]] = v
109
+ else:
110
+ cleaned_state[k] = v
111
+
112
+ loaded_count = 0
113
+
114
+ # 1. Bond expansion centers/width if present
115
+ if hasattr(model, "bond_expansion") and hasattr(model.bond_expansion, "rbf"):
116
+ rbf = model.bond_expansion.rbf
117
+ for k in ("bond_expansion.rbf.centers", "bond_expansion.centers"):
118
+ if k in cleaned_state and hasattr(rbf, "centers"):
119
+ rbf.centers.assign(cleaned_state[k].detach().cpu().numpy())
120
+ loaded_count += 1
121
+ for k in ("bond_expansion.rbf.width", "bond_expansion.width"):
122
+ if k in cleaned_state and hasattr(rbf, "width"):
123
+ rbf.width.assign(cleaned_state[k].detach().cpu().numpy())
124
+ loaded_count += 1
125
+
126
+ # 2. Embedding block
127
+ if hasattr(model, "embedding"):
128
+ emb = model.embedding
129
+ if hasattr(emb, "layer_node_embedding") and "embedding.layer_node_embedding.weight" in cleaned_state:
130
+ w = cleaned_state["embedding.layer_node_embedding.weight"].detach().cpu().numpy()
131
+ emb.layer_node_embedding.set_weights([w])
132
+ loaded_count += 1
133
+ if hasattr(emb, "emb") and "embedding.emb.weight" in cleaned_state:
134
+ w = cleaned_state["embedding.emb.weight"].detach().cpu().numpy()
135
+ emb.emb.set_weights([w])
136
+ loaded_count += 1
137
+
138
+ # 3. Traverse model sublayers and match with state keys
139
+ for name, sublayer in model.__dict__.items():
140
+ if isinstance(sublayer, keras.layers.Layer):
141
+ # Check for direct linear weight matches
142
+ w_key = f"{name}.weight"
143
+ b_key = f"{name}.bias"
144
+ if w_key in cleaned_state:
145
+ _assign_dense_weights(sublayer, cleaned_state[w_key], cleaned_state.get(b_key))
146
+ loaded_count += 1
147
+
148
+ if log:
149
+ print(f"Loaded {loaded_count} weight tensors into {model.__class__.__name__}")
150
+
151
+ return loaded_count
152
+
153
+
154
+ def load_model(name_or_path: str, **kwargs) -> keras.Model:
155
+ """Convenience factory to download/load and instantiate any MatGL model."""
156
+ from .megnet import MEGNet
157
+ from .m3gnet import M3GNet
158
+ from .tensornet import TensorNet
159
+ from .chgnet import CHGNet
160
+ from .so3net import SO3Net
161
+ from .grace import GRACE
162
+ from .qet import QET
163
+ from .wrappers import TransformedTargetModel
164
+
165
+ if os.path.exists(name_or_path) and os.path.isdir(name_or_path):
166
+ f_json = os.path.join(name_or_path, "model.json")
167
+ f_state = os.path.join(name_or_path, "state.pt")
168
+ else:
169
+ files = download_matgl_checkpoint(name_or_path, **kwargs)
170
+ f_json = files["model.json"]
171
+ f_state = files["state.pt"]
172
+
173
+ with open(f_json) as f:
174
+ meta = json.load(f)
175
+
176
+ cls_name = meta.get("@class")
177
+ init_kwargs = meta.get("kwargs", {})
178
+
179
+ # Handle TransformedTargetModel
180
+ if cls_name == "TransformedTargetModel":
181
+ inner_info = init_kwargs.get("model", {})
182
+ inner_cls = inner_info.get("@class", "MEGNet")
183
+ inner_args = inner_info.get("init_args", {})
184
+ transformer_info = init_kwargs.get("target_transformer", {})
185
+ mean = float(transformer_info.get("mean", 0.0))
186
+ std = float(transformer_info.get("std", 1.0))
187
+
188
+ if inner_cls == "MEGNet":
189
+ base_model = MEGNet(**{k: v for k, v in inner_args.items() if k in (
190
+ "dim_node_embedding", "dim_edge_embedding", "dim_state_embedding",
191
+ "nblocks", "cutoff"
192
+ )})
193
+ elif inner_cls == "M3GNet":
194
+ base_model = M3GNet(**{k: v for k, v in inner_args.items() if k in (
195
+ "dim_node_embedding", "dim_edge_embedding", "nblocks", "cutoff", "threebody_cutoff"
196
+ )})
197
+ elif inner_cls == "TensorNet":
198
+ base_model = TensorNet(**{k: v for k, v in inner_args.items() if k in (
199
+ "units", "nblocks", "num_rbf", "cutoff"
200
+ )})
201
+ elif inner_cls == "CHGNet":
202
+ base_model = CHGNet(**{k: v for k, v in inner_args.items() if k in (
203
+ "dim_atom_embedding", "dim_bond_embedding", "cutoff", "threebody_cutoff"
204
+ )})
205
+ else:
206
+ base_model = MEGNet()
207
+
208
+ load_matgl_weights(base_model, f_state)
209
+ return TransformedTargetModel(model=base_model, mean=mean, std=std)
210
+
211
+ elif cls_name == "MEGNet":
212
+ model = MEGNet()
213
+ elif cls_name == "M3GNet":
214
+ model = M3GNet()
215
+ elif cls_name == "TensorNet":
216
+ model = TensorNet()
217
+ elif cls_name == "CHGNet":
218
+ model = CHGNet()
219
+ elif cls_name == "SO3Net":
220
+ model = SO3Net()
221
+ elif cls_name == "GRACE":
222
+ model = GRACE()
223
+ elif cls_name == "QET":
224
+ model = QET()
225
+ else:
226
+ model = MEGNet()
227
+
228
+ load_matgl_weights(model, f_state)
229
+ return model
230
+
@@ -0,0 +1,462 @@
1
+ """Multi-backend Keras 3 implementation of M3GNet."""
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, GatedMLP, EmbeddingBlock, scatter_add, scatter_mean, infer_num_graphs
11
+ from .basis import (
12
+ BondExpansion,
13
+ SphericalBesselWithHarmonics,
14
+ compute_pair_vector_and_distance,
15
+ compute_theta_and_phi,
16
+ cosine_cutoff,
17
+ )
18
+ from .readout import WeightedAtomReadOut, ReduceReadOut, Set2SetReadOut
19
+
20
+
21
+ class ThreeBodyInteractions(layers.Layer):
22
+ """Three-body bond angular update using directed line-graph message passing.
23
+
24
+ Example:
25
+ ```python
26
+ import numpy as np
27
+ import keras
28
+ from k3_node.models import ThreeBodyInteractions
29
+
30
+ edge_index = np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]) # 4 atoms, 8 bonds
31
+ node_feat = np.random.rand(4, 16).astype("float32")
32
+ edge_feat = np.random.rand(8, 16).astype("float32")
33
+ line_edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]]) # bond pairs forming 4 angles
34
+ three_basis = np.random.rand(4, 9).astype("float32") # angular basis of each triplet
35
+ three_cutoff = np.ones(8, dtype="float32") # smooth cutoff weight of each bond
36
+
37
+ layer = ThreeBodyInteractions(
38
+ update_network_atom=keras.layers.Dense(9, activation="sigmoid"),
39
+ update_network_bond=keras.layers.Dense(16),
40
+ )
41
+ print(tuple(layer(edge_index, line_edge_index, three_basis, three_cutoff, node_feat, edge_feat).shape)) # (8, 16)
42
+ ```
43
+ """
44
+
45
+ def __init__(
46
+ self,
47
+ update_network_atom: layers.Layer,
48
+ update_network_bond: layers.Layer,
49
+ **kwargs,
50
+ ):
51
+ super().__init__(**kwargs)
52
+ self.update_network_atom = update_network_atom
53
+ self.update_network_bond = update_network_bond
54
+
55
+ def call(
56
+ self,
57
+ edge_index,
58
+ line_edge_index,
59
+ three_basis,
60
+ three_cutoff,
61
+ node_feat,
62
+ edge_feat,
63
+ ):
64
+ src_bond = ops.cast(line_edge_index[0], "int32")
65
+ dst_bond = ops.cast(line_edge_index[1], "int32")
66
+ edge_dst_atom = ops.cast(edge_index[1], "int32")
67
+ num_bonds = ops.shape(edge_feat)[0]
68
+
69
+ # Destination atom of destination bond
70
+ end_atom_idx = ops.take(edge_dst_atom, dst_bond, axis=0)
71
+ updated_atoms = self.update_network_atom(node_feat)
72
+ end_atom_features = ops.take(updated_atoms, end_atom_idx, axis=0)
73
+
74
+ basis = three_basis * end_atom_features
75
+ weights = ops.take(three_cutoff, src_bond, axis=0) * ops.take(three_cutoff, dst_bond, axis=0)
76
+ basis = basis * ops.expand_dims(weights, axis=-1)
77
+
78
+ new_bonds = scatter_add(basis, src_bond, num_segments=num_bonds)
79
+ return edge_feat + self.update_network_bond(new_bonds)
80
+
81
+
82
+ class M3GNetGraphConv(layers.Layer):
83
+ """M3GNet graph convolution layer: two-body edge and node updates.
84
+
85
+ Example:
86
+ ```python
87
+ import numpy as np
88
+ from k3_node.models import M3GNetGraphConv
89
+
90
+ edge_index = np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]) # 4 atoms, 8 bonds
91
+ node_feat = np.random.rand(4, 16).astype("float32")
92
+ edge_feat = np.random.rand(8, 16).astype("float32")
93
+ rbf = np.random.rand(8, 9).astype("float32") # radial basis of each bond
94
+
95
+ conv = M3GNetGraphConv(degree=9, edge_dims=[48, 16, 16], node_dims=[48, 16, 16])
96
+ edge_out, node_out, state_out = conv(edge_index, edge_feat, node_feat, None, rbf, num_nodes=4)
97
+ print(tuple(edge_out.shape), tuple(node_out.shape)) # (8, 16) (4, 16)
98
+ ```
99
+ """
100
+
101
+ def __init__(
102
+ self,
103
+ degree: int,
104
+ edge_dims: Sequence[int],
105
+ node_dims: Sequence[int],
106
+ state_dims: Optional[Sequence[int]] = None,
107
+ include_state: bool = False,
108
+ activation: str = "swish",
109
+ **kwargs,
110
+ ):
111
+ super().__init__(**kwargs)
112
+ self.include_state = include_state
113
+ self.edge_update_func = GatedMLP(in_feats=edge_dims[0], dims=edge_dims[1:])
114
+ self.edge_weight_func = layers.Dense(edge_dims[-1], use_bias=False)
115
+
116
+ self.node_update_func = GatedMLP(in_feats=node_dims[0], dims=node_dims[1:])
117
+ self.node_weight_func = layers.Dense(node_dims[-1], use_bias=False)
118
+
119
+ if include_state and state_dims is not None:
120
+ self.state_update_func = MLP(state_dims, activation=activation, activate_last=True)
121
+ else:
122
+ self.state_update_func = None
123
+
124
+ def call(
125
+ self,
126
+ edge_index,
127
+ edge_feat,
128
+ node_feat,
129
+ state_feat,
130
+ rbf,
131
+ batch=None,
132
+ edge_batch=None,
133
+ num_nodes: Optional[int] = None,
134
+ num_graphs: Optional[int] = None,
135
+ ):
136
+ src = ops.cast(edge_index[0], "int32")
137
+ dst = ops.cast(edge_index[1], "int32")
138
+
139
+ vi = ops.take(node_feat, src, axis=0)
140
+ vj = ops.take(node_feat, dst, axis=0)
141
+
142
+ if self.include_state and state_feat is not None:
143
+ if edge_batch is not None:
144
+ u_edge = ops.take(state_feat, edge_batch, axis=0)
145
+ else:
146
+ num_e = ops.shape(edge_feat)[0]
147
+ u_edge = ops.broadcast_to(state_feat, (num_e, ops.shape(state_feat)[-1]))
148
+ edge_inputs = ops.concatenate([vi, vj, edge_feat, u_edge], axis=-1)
149
+ else:
150
+ edge_inputs = ops.concatenate([vi, vj, edge_feat], axis=-1)
151
+
152
+ # 1. Edge update
153
+ edge_update = self.edge_update_func(edge_inputs) * self.edge_weight_func(rbf)
154
+ edge_feat_new = edge_feat + edge_update
155
+
156
+ # 2. Node update
157
+ node_update = self.node_update_func(edge_inputs) * self.node_weight_func(rbf)
158
+ node_update_sum = scatter_add(node_update, src, num_segments=num_nodes)
159
+ node_feat_new = node_feat + node_update_sum
160
+
161
+ # 3. State update
162
+ state_feat_new = state_feat
163
+ if self.include_state and self.state_update_func is not None and state_feat is not None:
164
+ if batch is not None:
165
+ uv = scatter_mean(node_feat_new, batch, num_segments=num_graphs)
166
+ else:
167
+ uv = ops.mean(node_feat_new, axis=0, keepdims=True)
168
+ state_inputs = ops.concatenate([state_feat, uv], axis=-1)
169
+ state_feat_new = self.state_update_func(state_inputs)
170
+
171
+ return edge_feat_new, node_feat_new, state_feat_new
172
+
173
+
174
+ class M3GNetBlock(layers.Layer):
175
+ """M3GNet block wrapping M3GNetGraphConv with optional dropout.
176
+
177
+ Example:
178
+ ```python
179
+ import numpy as np
180
+ from k3_node.models import M3GNetBlock
181
+
182
+ edge_index = np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]) # 4 atoms, 8 bonds
183
+ node_feat = np.random.rand(4, 16).astype("float32")
184
+ edge_feat = np.random.rand(8, 16).astype("float32")
185
+ rbf = np.random.rand(8, 9).astype("float32") # radial basis of each bond
186
+
187
+ block = M3GNetBlock(degree=9, conv_hiddens=[16], dim_node_feats=16, dim_edge_feats=16)
188
+ edge_out, node_out, state_out = block(edge_index, edge_feat, node_feat, None, rbf, num_nodes=4)
189
+ print(tuple(edge_out.shape), tuple(node_out.shape)) # (8, 16) (4, 16)
190
+ ```
191
+ """
192
+
193
+ def __init__(
194
+ self,
195
+ degree: int,
196
+ conv_hiddens: Sequence[int],
197
+ dim_node_feats: int,
198
+ dim_edge_feats: int,
199
+ dim_state_feats: int = 0,
200
+ include_state: bool = False,
201
+ activation: str = "swish",
202
+ dropout: float = 0.0,
203
+ **kwargs,
204
+ ):
205
+ super().__init__(**kwargs)
206
+ self.include_state = include_state
207
+ edge_in = 2 * dim_node_feats + dim_edge_feats + (dim_state_feats if include_state else 0)
208
+ node_in = 2 * dim_node_feats + dim_edge_feats + (dim_state_feats if include_state else 0)
209
+ state_in = dim_state_feats + dim_node_feats if include_state else 0
210
+
211
+ self.conv = M3GNetGraphConv(
212
+ degree=degree,
213
+ edge_dims=[edge_in, *conv_hiddens, dim_edge_feats],
214
+ node_dims=[node_in, *conv_hiddens, dim_node_feats],
215
+ state_dims=[state_in, *conv_hiddens, dim_state_feats] if include_state else None,
216
+ include_state=include_state,
217
+ activation=activation,
218
+ )
219
+ self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
220
+
221
+ def call(
222
+ self,
223
+ edge_index,
224
+ edge_feat,
225
+ node_feat,
226
+ state_feat,
227
+ rbf,
228
+ batch=None,
229
+ edge_batch=None,
230
+ num_nodes: Optional[int] = None,
231
+ num_graphs: Optional[int] = None,
232
+ training=None,
233
+ ):
234
+ edge_feat, node_feat, state_feat = self.conv(
235
+ edge_index, edge_feat, node_feat, state_feat, rbf,
236
+ batch=batch, edge_batch=edge_batch, num_nodes=num_nodes, num_graphs=num_graphs
237
+ )
238
+ if self.dropout is not None:
239
+ edge_feat = self.dropout(edge_feat, training=training)
240
+ node_feat = self.dropout(node_feat, training=training)
241
+ return edge_feat, node_feat, state_feat
242
+
243
+
244
+ class M3GNet(keras.Model):
245
+ """M3GNet materials potential model supporting 3-body angles and multibackend training.
246
+
247
+ Example:
248
+ ```python
249
+ import numpy as np
250
+ from k3_node.models import M3GNet
251
+
252
+ # A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
253
+ structure = {
254
+ "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"),
255
+ "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
256
+ "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), # bond pairs forming angles
257
+ "node_type": np.array([6, 8, 1, 6]), # atomic numbers
258
+ "batch": np.zeros(4, dtype="int32"), # all atoms belong to structure 0
259
+ "state_attr": np.zeros((1, 2), dtype="float32"), # global state features
260
+ }
261
+
262
+ model = M3GNet(dim_node_embedding=16, dim_edge_embedding=16, nblocks=2, units=16, max_n=3, max_l=3)
263
+ energy = model(structure) # predicted property (e.g. energy) of the structure
264
+ print(tuple(energy.shape)) # (1,)
265
+ ```
266
+ """
267
+
268
+ def __init__(
269
+ self,
270
+ dim_node_embedding: int = 64,
271
+ dim_edge_embedding: int = 64,
272
+ dim_state_embedding: int = 0,
273
+ ntypes_state: Optional[int] = None,
274
+ max_n: int = 3,
275
+ max_l: int = 3,
276
+ nblocks: int = 3,
277
+ rbf_type: Literal["Gaussian", "SphericalBessel"] = "SphericalBessel",
278
+ is_intensive: bool = True,
279
+ readout_type: Literal["set2set", "weighted_atom", "reduce_atom"] = "weighted_atom",
280
+ cutoff: float = 5.0,
281
+ threebody_cutoff: float = 4.0,
282
+ units: int = 64,
283
+ ntargets: int = 1,
284
+ include_state: bool = False,
285
+ activation_type: str = "swish",
286
+ dropout: float = 0.0,
287
+ ntypes_node: int = 95,
288
+ **kwargs,
289
+ ):
290
+ super().__init__(**kwargs)
291
+ self.cutoff = cutoff
292
+ self.threebody_cutoff = threebody_cutoff
293
+ self.include_state = include_state
294
+ self.is_intensive = is_intensive
295
+
296
+ self.bond_expansion = BondExpansion(
297
+ max_l=max_l,
298
+ max_n=max_n,
299
+ cutoff=cutoff,
300
+ rbf_type=rbf_type,
301
+ smooth=False,
302
+ )
303
+
304
+ degree_rbf = max_n * max_l if rbf_type.lower() == "sphericalbessel" else 100
305
+ self.embedding = EmbeddingBlock(
306
+ degree_rbf=degree_rbf,
307
+ dim_node_embedding=dim_node_embedding,
308
+ dim_edge_embedding=dim_edge_embedding,
309
+ ntypes_node=ntypes_node,
310
+ ntypes_state=ntypes_state,
311
+ include_state=include_state,
312
+ dim_state_embedding=dim_state_embedding if include_state else None,
313
+ activation=activation_type,
314
+ )
315
+
316
+ self.sbf_shf = SphericalBesselWithHarmonics(
317
+ max_n=max_n,
318
+ max_l=max_l,
319
+ cutoff=threebody_cutoff,
320
+ )
321
+ degree_3body = max_n * max_l
322
+
323
+ self.three_body_interactions = []
324
+ self.blocks = []
325
+ for _ in range(nblocks):
326
+ self.three_body_interactions.append(
327
+ ThreeBodyInteractions(
328
+ update_network_atom=MLP([dim_node_embedding, degree_3body], activation=activation_type, activate_last=False),
329
+ update_network_bond=GatedMLP(in_feats=degree_3body, dims=[dim_edge_embedding], activate_last=False),
330
+ )
331
+ )
332
+ self.blocks.append(
333
+ M3GNetBlock(
334
+ degree=degree_rbf,
335
+ conv_hiddens=[units, units],
336
+ dim_node_feats=dim_node_embedding,
337
+ dim_edge_feats=dim_edge_embedding,
338
+ dim_state_feats=dim_state_embedding if include_state else 0,
339
+ include_state=include_state,
340
+ activation=activation_type,
341
+ dropout=dropout,
342
+ )
343
+ )
344
+
345
+ if readout_type == "weighted_atom":
346
+ self.readout = WeightedAtomReadOut(dim_node_embedding, dims=[units, units, ntargets], activation=activation_type)
347
+ elif readout_type == "set2set":
348
+ self.readout = Set2SetReadOut(dim_node_embedding)
349
+ self.final_mlp = MLP([2 * dim_node_embedding, units, ntargets], activation=activation_type, activate_last=False)
350
+ else:
351
+ self.readout = ReduceReadOut(op="mean")
352
+ self.final_mlp = MLP([dim_node_embedding, units, ntargets], activation=activation_type, activate_last=False)
353
+
354
+ def _unpack_inputs(self, inputs):
355
+ if isinstance(inputs, dict):
356
+ pos = inputs.get("pos")
357
+ edge_index = inputs.get("edge_index")
358
+ node_type = inputs.get("node_type", inputs.get("z"))
359
+ line_edge_index = inputs.get("line_edge_index", None)
360
+ state_attr = inputs.get("state_attr", None)
361
+ pbc_offshift = inputs.get("pbc_offshift", None)
362
+ batch = inputs.get("batch", None)
363
+ num_graphs = inputs.get("num_graphs", None)
364
+ return pos, edge_index, node_type, line_edge_index, state_attr, pbc_offshift, batch, num_graphs
365
+ elif isinstance(inputs, (tuple, list)):
366
+ pos = inputs[0]
367
+ edge_index = inputs[1]
368
+ node_type = inputs[2]
369
+ line_edge_index = inputs[3] if len(inputs) > 3 else None
370
+ state_attr = inputs[4] if len(inputs) > 4 else None
371
+ pbc_offshift = inputs[5] if len(inputs) > 5 else None
372
+ batch = inputs[6] if len(inputs) > 6 else None
373
+ num_graphs = inputs[7] if len(inputs) > 7 else None
374
+ return pos, edge_index, node_type, line_edge_index, state_attr, pbc_offshift, batch, num_graphs
375
+ return inputs, None, None, None, None, None, None, None
376
+
377
+ def call(
378
+ self,
379
+ inputs,
380
+ edge_index=None,
381
+ node_type=None,
382
+ line_edge_index=None,
383
+ state_attr=None,
384
+ pbc_offshift=None,
385
+ batch=None,
386
+ num_graphs=None,
387
+ training=None,
388
+ ):
389
+ if edge_index is None:
390
+ (
391
+ pos,
392
+ edge_index,
393
+ node_type_in,
394
+ line_edge_index_in,
395
+ state_attr_in,
396
+ pbc_offshift_in,
397
+ batch_in,
398
+ num_graphs_in,
399
+ ) = self._unpack_inputs(inputs)
400
+ if node_type is None:
401
+ node_type = node_type_in
402
+ if line_edge_index is None:
403
+ line_edge_index = line_edge_index_in
404
+ if state_attr is None:
405
+ state_attr = state_attr_in
406
+ if pbc_offshift is None:
407
+ pbc_offshift = pbc_offshift_in
408
+ if batch is None:
409
+ batch = batch_in
410
+ if num_graphs is None:
411
+ num_graphs = num_graphs_in
412
+ else:
413
+ pos = inputs
414
+
415
+ num_nodes = ops.shape(pos)[0]
416
+ if batch is None:
417
+ batch = ops.zeros((num_nodes,), dtype="int32")
418
+ else:
419
+ batch = ops.cast(batch, "int32")
420
+ n_graphs = infer_num_graphs(batch=batch, num_graphs=num_graphs, state_attr=state_attr)
421
+
422
+ src = ops.cast(edge_index[0], "int32")
423
+ edge_batch = ops.take(batch, src, axis=0)
424
+
425
+ # 1. 2-body pair distances and expansion
426
+ _, bond_dists = compute_pair_vector_and_distance(pos, edge_index, pbc_offshift)
427
+ edge_attr = self.bond_expansion(bond_dists)
428
+
429
+ # 2. Embeddings
430
+ node_feat, edge_feat, state_feat = self.embedding(node_type, edge_attr, state_attr)
431
+
432
+ # 3. 3-body expansion if line_edge_index is available
433
+ three_cutoff = cosine_cutoff(bond_dists, self.threebody_cutoff)
434
+ if line_edge_index is not None and ops.shape(line_edge_index)[1] > 0:
435
+ theta, phi = compute_theta_and_phi(pos, edge_index, line_edge_index, pbc_offshift)
436
+ src_bonds = ops.cast(line_edge_index[0], "int32")
437
+ r_triplets = ops.take(bond_dists, src_bonds, axis=0)
438
+ three_basis = self.sbf_shf(r_triplets, theta, phi)
439
+ else:
440
+ three_basis = None
441
+
442
+ # 4. Message passing blocks
443
+ for i in range(len(self.blocks)):
444
+ if three_basis is not None and line_edge_index is not None:
445
+ edge_feat = self.three_body_interactions[i](
446
+ edge_index, line_edge_index, three_basis, three_cutoff, node_feat, edge_feat
447
+ )
448
+ edge_feat, node_feat, state_feat = self.blocks[i](
449
+ edge_index, edge_feat, node_feat, state_feat, edge_attr,
450
+ batch=batch, edge_batch=edge_batch, num_nodes=num_nodes, num_graphs=n_graphs,
451
+ training=training,
452
+ )
453
+
454
+ # 5. Readout
455
+ if isinstance(self.readout, WeightedAtomReadOut):
456
+ output = self.readout(node_feat, batch=batch, num_graphs=n_graphs)
457
+ else:
458
+ pooled = self.readout(node_feat, batch=batch, num_graphs=n_graphs)
459
+ output = self.final_mlp(pooled)
460
+
461
+ return ops.squeeze(output, axis=-1)
462
+