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,381 @@
1
+ """Multi-backend Keras 3 implementation of TensorNet."""
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 (
11
+ MLP,
12
+ vector_to_skewtensor,
13
+ vector_to_symtensor,
14
+ decompose_tensor,
15
+ tensor_norm,
16
+ scatter_add,
17
+ infer_num_graphs,
18
+ )
19
+ from .basis import (
20
+ BondExpansion,
21
+ RadialBesselFunction,
22
+ compute_pair_vector_and_distance,
23
+ cosine_cutoff,
24
+ )
25
+ from .readout import WeightedAtomReadOut, ReduceReadOut
26
+
27
+
28
+ class TensorEmbedding(layers.Layer):
29
+ """Embeds node types and Cartesian pair vectors into rank-2 tensors [num_nodes, units, 3, 3].
30
+
31
+ Example:
32
+ ```python
33
+ import numpy as np
34
+ from k3_node.models import TensorEmbedding
35
+
36
+ z = np.array([6, 8, 1, 6]) # atomic numbers
37
+ edge_index = np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]])
38
+ edge_attr = np.random.rand(8, 16).astype("float32") # radial basis of each bond
39
+ edge_weight = np.random.rand(8).astype("float32") * 3.0 # bond lengths
40
+ vec = np.random.rand(8, 3).astype("float32") # bond vectors
41
+ layer = TensorEmbedding(units=16, degree_rbf=16)
42
+ print(tuple(layer(z, edge_index, edge_attr, edge_weight, vec).shape)) # (4, 16, 3, 3): a 3x3 tensor per channel
43
+ ```
44
+ """
45
+
46
+ def __init__(
47
+ self,
48
+ units: int,
49
+ degree_rbf: int,
50
+ ntypes_node: int = 95,
51
+ cutoff: float = 5.0,
52
+ activation: str = "swish",
53
+ **kwargs,
54
+ ):
55
+ super().__init__(**kwargs)
56
+ self.units = units
57
+ self.cutoff = cutoff
58
+
59
+ self.distance_proj1 = layers.Dense(units, use_bias=True)
60
+ self.distance_proj2 = layers.Dense(units, use_bias=True)
61
+ self.distance_proj3 = layers.Dense(units, use_bias=True)
62
+
63
+ self.emb = layers.Embedding(ntypes_node, units)
64
+ self.emb2 = layers.Dense(units, use_bias=True)
65
+
66
+ self.linears_tensor = [layers.Dense(units, use_bias=False) for _ in range(3)]
67
+ self.linears_scalar = [
68
+ layers.Dense(2 * units, use_bias=True),
69
+ layers.Dense(3 * units, use_bias=True),
70
+ ]
71
+ self.init_norm = layers.LayerNormalization(axis=-1)
72
+
73
+ def call(self, z, edge_index, edge_attr, edge_weight, vec):
74
+ src = ops.cast(edge_index[0], "int32")
75
+ dst = ops.cast(edge_index[1], "int32")
76
+ num_nodes = ops.shape(z)[0]
77
+
78
+ # Normalized pair vectors
79
+ vec_norm = vec / ops.maximum(ops.expand_dims(edge_weight, axis=-1), 1e-7)
80
+
81
+ # Distance projections: [num_edges, units]
82
+ C = ops.expand_dims(cosine_cutoff(edge_weight, self.cutoff), axis=-1)
83
+ f_I = self.distance_proj1(edge_attr) * C
84
+ f_A = self.distance_proj2(edge_attr) * C
85
+ f_S = self.distance_proj3(edge_attr) * C
86
+
87
+ # Geometric tensors: [num_edges, 3, 3]
88
+ I_mat = ops.expand_dims(ops.eye(3, dtype=vec.dtype), axis=0) # [1, 3, 3]
89
+ A_mat = vector_to_skewtensor(vec_norm) # [num_edges, 3, 3]
90
+ S_mat = vector_to_symtensor(vec_norm) # [num_edges, 3, 3]
91
+
92
+ # Expand to units dimension: [num_edges, units, 3, 3]
93
+ Iij = ops.expand_dims(ops.expand_dims(f_I, axis=-1), axis=-1) * ops.expand_dims(I_mat, axis=1)
94
+ Aij = ops.expand_dims(ops.expand_dims(f_A, axis=-1), axis=-1) * ops.expand_dims(A_mat, axis=1)
95
+ Sij = ops.expand_dims(ops.expand_dims(f_S, axis=-1), axis=-1) * ops.expand_dims(S_mat, axis=1)
96
+
97
+ # Node chemical embeddings
98
+ node_emb = self.emb(ops.cast(z, "int32"))
99
+ vi = ops.take(node_emb, src, axis=0)
100
+ vj = ops.take(node_emb, dst, axis=0)
101
+ zij = ops.concatenate([vi, vj], axis=-1)
102
+ Zij = ops.expand_dims(ops.expand_dims(self.emb2(zij), axis=-1), axis=-1)
103
+
104
+ scalars_msg = Zij * Iij
105
+ skew_msg = Zij * Aij
106
+ traceless_msg = Zij * Sij
107
+
108
+ scalars = scatter_add(scalars_msg, src, num_segments=num_nodes)
109
+ skew = scatter_add(skew_msg, src, num_segments=num_nodes)
110
+ traceless = scatter_add(traceless_msg, src, num_segments=num_nodes)
111
+
112
+ # Apply tensor linear transformations: transpose to apply Dense along units axis
113
+ s_t = ops.transpose(scalars, [0, 2, 3, 1])
114
+ a_t = ops.transpose(skew, [0, 2, 3, 1])
115
+ tr_t = ops.transpose(traceless, [0, 2, 3, 1])
116
+
117
+ scalars = ops.transpose(self.linears_tensor[0](s_t), [0, 3, 1, 2])
118
+ skew = ops.transpose(self.linears_tensor[1](a_t), [0, 3, 1, 2])
119
+ traceless = ops.transpose(self.linears_tensor[2](tr_t), [0, 3, 1, 2])
120
+
121
+ # Node invariant scalar feature
122
+ s_norm = self.init_norm(tensor_norm(scalars))
123
+ s_norm = ops.silu(self.linears_scalar[0](s_norm))
124
+ s_norm = self.linears_scalar[1](s_norm)
125
+
126
+ f_I_node = s_norm[..., :self.units]
127
+ f_A_node = s_norm[..., self.units:2 * self.units]
128
+ f_S_node = s_norm[..., 2 * self.units:]
129
+
130
+ scalars = ops.expand_dims(ops.expand_dims(f_I_node, axis=-1), axis=-1) * scalars
131
+ skew = ops.expand_dims(ops.expand_dims(f_A_node, axis=-1), axis=-1) * skew
132
+ traceless = ops.expand_dims(ops.expand_dims(f_S_node, axis=-1), axis=-1) * traceless
133
+
134
+ X = scalars + skew + traceless
135
+ return X
136
+
137
+
138
+ class TensorNetInteraction(layers.Layer):
139
+ """Equivariant Cartesian tensor message passing interaction layer.
140
+
141
+ Example:
142
+ ```python
143
+ import numpy as np
144
+ from k3_node.models import TensorNetInteraction
145
+
146
+ edge_index = np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]])
147
+ edge_weight = np.random.rand(8).astype("float32") * 3.0
148
+ edge_attr = np.random.rand(8, 16).astype("float32")
149
+ X = np.random.rand(4, 16, 3, 3).astype("float32") # per-atom tensor features
150
+ layer = TensorNetInteraction(num_rbf=16, units=16)
151
+ print(tuple(layer(edge_index, edge_weight, edge_attr, X).shape)) # (4, 16, 3, 3)
152
+ ```
153
+ """
154
+
155
+ def __init__(
156
+ self,
157
+ num_rbf: int,
158
+ units: int,
159
+ cutoff: float = 5.0,
160
+ activation: str = "swish",
161
+ **kwargs,
162
+ ):
163
+ super().__init__(**kwargs)
164
+ self.num_rbf = num_rbf
165
+ self.units = units
166
+ self.cutoff = cutoff
167
+
168
+ self.linears_scalar = [
169
+ layers.Dense(units, use_bias=True),
170
+ layers.Dense(2 * units, use_bias=True),
171
+ layers.Dense(3 * units, use_bias=True),
172
+ ]
173
+ self.linears_tensor = [layers.Dense(units, use_bias=False) for _ in range(6)]
174
+
175
+ def call(self, edge_index, edge_weight, edge_attr, X):
176
+ src = ops.cast(edge_index[0], "int32")
177
+ dst = ops.cast(edge_index[1], "int32")
178
+ num_nodes = ops.shape(X)[0]
179
+
180
+ # Process edge attributes
181
+ C = ops.expand_dims(cosine_cutoff(edge_weight, self.cutoff), axis=-1)
182
+ h_edge = edge_attr
183
+ for linear in self.linears_scalar:
184
+ h_edge = ops.silu(linear(h_edge))
185
+ edge_attr_processed = h_edge * C
186
+
187
+ f_I = edge_attr_processed[..., :self.units]
188
+ f_A = edge_attr_processed[..., :self.units]
189
+ f_S = edge_attr_processed[..., :self.units]
190
+
191
+ # Normalize input tensor
192
+ X_norm = ops.expand_dims(ops.expand_dims(tensor_norm(X) + 1.0, axis=-1), axis=-1)
193
+ X_normalized = X / X_norm
194
+
195
+ scalars, skew, traceless = decompose_tensor(X_normalized)
196
+
197
+ # Tensor linears (0, 1, 2)
198
+ s_t = ops.transpose(scalars, [0, 2, 3, 1])
199
+ a_t = ops.transpose(skew, [0, 2, 3, 1])
200
+ tr_t = ops.transpose(traceless, [0, 2, 3, 1])
201
+
202
+ scalars = ops.transpose(self.linears_tensor[0](s_t), [0, 3, 1, 2])
203
+ skew = ops.transpose(self.linears_tensor[1](a_t), [0, 3, 1, 2])
204
+ traceless = ops.transpose(self.linears_tensor[2](tr_t), [0, 3, 1, 2])
205
+
206
+ # Gather node features for edges
207
+ sc_j = ops.take(scalars, dst, axis=0)
208
+ sk_j = ops.take(skew, dst, axis=0)
209
+ tr_j = ops.take(traceless, dst, axis=0)
210
+
211
+ # Modulate by radial features
212
+ msg_s = ops.expand_dims(ops.expand_dims(f_I, axis=-1), axis=-1) * sc_j
213
+ msg_a = ops.expand_dims(ops.expand_dims(f_A, axis=-1), axis=-1) * sk_j
214
+ msg_tr = ops.expand_dims(ops.expand_dims(f_S, axis=-1), axis=-1) * tr_j
215
+
216
+ msg = msg_s + msg_a + msg_tr
217
+ X_update = scatter_add(msg, src, num_segments=num_nodes)
218
+
219
+ # Tensor linears (3, 4, 5)
220
+ s2, a2, tr2 = decompose_tensor(X_update)
221
+ s2_t = ops.transpose(s2, [0, 2, 3, 1])
222
+ a2_t = ops.transpose(a2, [0, 2, 3, 1])
223
+ tr2_t = ops.transpose(tr2, [0, 2, 3, 1])
224
+
225
+ s2 = ops.transpose(self.linears_tensor[3](s2_t), [0, 3, 1, 2])
226
+ a2 = ops.transpose(self.linears_tensor[4](a2_t), [0, 3, 1, 2])
227
+ tr2 = ops.transpose(self.linears_tensor[5](tr2_t), [0, 3, 1, 2])
228
+
229
+ return X + s2 + a2 + tr2
230
+
231
+
232
+ class TensorNet(keras.Model):
233
+ """Cartesian tensor-based equivariant GNN for molecular and crystal potentials.
234
+
235
+ Example:
236
+ ```python
237
+ import numpy as np
238
+ from k3_node.models import TensorNet
239
+
240
+ # A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
241
+ structure = {
242
+ "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"),
243
+ "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
244
+ "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), # bond pairs forming angles
245
+ "node_type": np.array([6, 8, 1, 6]), # atomic numbers
246
+ "batch": np.zeros(4, dtype="int32"), # all atoms belong to structure 0
247
+ "state_attr": np.zeros((1, 2), dtype="float32"), # global state features
248
+ }
249
+
250
+ model = TensorNet(units=16, nblocks=2, num_rbf=16)
251
+ energy = model(structure) # predicted property (e.g. energy) of the structure
252
+ print(tuple(energy.shape)) # (1,)
253
+ ```
254
+ """
255
+
256
+ def __init__(
257
+ self,
258
+ units: int = 64,
259
+ nblocks: int = 2,
260
+ num_rbf: int = 32,
261
+ cutoff: float = 5.0,
262
+ rbf_type: Literal["Gaussian", "SphericalBessel"] = "Gaussian",
263
+ ntypes_node: int = 95,
264
+ ntargets: int = 1,
265
+ readout_type: Literal["weighted_atom", "reduce_atom"] = "weighted_atom",
266
+ activation_type: str = "swish",
267
+ **kwargs,
268
+ ):
269
+ super().__init__(**kwargs)
270
+ self.units = units
271
+ self.cutoff = cutoff
272
+
273
+ if rbf_type.lower() == "gaussian":
274
+ self.bond_expansion = BondExpansion(
275
+ rbf_type="Gaussian",
276
+ initial=0.0,
277
+ final=cutoff,
278
+ num_centers=num_rbf,
279
+ )
280
+ else:
281
+ self.bond_expansion = RadialBesselFunction(max_n=num_rbf, cutoff=cutoff)
282
+
283
+ self.embedding = TensorEmbedding(
284
+ units=units,
285
+ degree_rbf=num_rbf,
286
+ ntypes_node=ntypes_node,
287
+ cutoff=cutoff,
288
+ activation=activation_type,
289
+ )
290
+
291
+ self.interactions = [
292
+ TensorNetInteraction(num_rbf=num_rbf, units=units, cutoff=cutoff, activation=activation_type)
293
+ for _ in range(nblocks)
294
+ ]
295
+
296
+ self.out_norm = layers.LayerNormalization(axis=-1)
297
+ self.node_proj = MLP([3 * units, units, units], activation=activation_type, activate_last=True)
298
+
299
+ if readout_type == "weighted_atom":
300
+ self.readout = WeightedAtomReadOut(units, dims=[units, ntargets], activation=activation_type)
301
+ else:
302
+ self.readout = ReduceReadOut(op="mean")
303
+ self.final_mlp = MLP([units, ntargets], activation=activation_type, activate_last=False)
304
+
305
+ def _unpack_inputs(self, inputs):
306
+ if isinstance(inputs, dict):
307
+ pos = inputs.get("pos")
308
+ edge_index = inputs.get("edge_index")
309
+ node_type = inputs.get("node_type", inputs.get("z"))
310
+ pbc_offshift = inputs.get("pbc_offshift", None)
311
+ batch = inputs.get("batch", None)
312
+ num_graphs = inputs.get("num_graphs", None)
313
+ state_attr = inputs.get("state_attr", None)
314
+ return pos, edge_index, node_type, pbc_offshift, batch, num_graphs, state_attr
315
+ elif isinstance(inputs, (tuple, list)):
316
+ pos = inputs[0]
317
+ edge_index = inputs[1]
318
+ node_type = inputs[2]
319
+ pbc_offshift = inputs[3] if len(inputs) > 3 else None
320
+ batch = inputs[4] if len(inputs) > 4 else None
321
+ num_graphs = inputs[5] if len(inputs) > 5 else None
322
+ state_attr = inputs[6] if len(inputs) > 6 else None
323
+ return pos, edge_index, node_type, pbc_offshift, batch, num_graphs, state_attr
324
+ return inputs, None, None, None, None, None, None
325
+
326
+ def call(self, inputs, edge_index=None, node_type=None, pbc_offshift=None, batch=None, num_graphs=None, state_attr=None):
327
+ if edge_index is None:
328
+ (
329
+ pos,
330
+ edge_index,
331
+ node_type_in,
332
+ pbc_offshift_in,
333
+ batch_in,
334
+ num_graphs_in,
335
+ state_attr_in,
336
+ ) = self._unpack_inputs(inputs)
337
+ if node_type is None:
338
+ node_type = node_type_in
339
+ if pbc_offshift is None:
340
+ pbc_offshift = pbc_offshift_in
341
+ if batch is None:
342
+ batch = batch_in
343
+ if num_graphs is None:
344
+ num_graphs = num_graphs_in
345
+ if state_attr is None:
346
+ state_attr = state_attr_in
347
+ else:
348
+ pos = inputs
349
+
350
+ num_nodes = ops.shape(pos)[0]
351
+ if batch is None:
352
+ batch = ops.zeros((num_nodes,), dtype="int32")
353
+ else:
354
+ batch = ops.cast(batch, "int32")
355
+ n_graphs = infer_num_graphs(batch=batch, num_graphs=num_graphs, state_attr=state_attr)
356
+
357
+ vec, bond_dists = compute_pair_vector_and_distance(pos, edge_index, pbc_offshift)
358
+ edge_attr = self.bond_expansion(bond_dists)
359
+
360
+ X = self.embedding(node_type, edge_index, edge_attr, bond_dists, vec)
361
+
362
+ for interaction in self.interactions:
363
+ X = interaction(edge_index, bond_dists, edge_attr, X)
364
+
365
+ # Decompose into irreducible norms
366
+ scalars, skew, traceless = decompose_tensor(X)
367
+ norm_s = tensor_norm(scalars)
368
+ norm_a = tensor_norm(skew)
369
+ norm_tr = tensor_norm(traceless)
370
+
371
+ norms = ops.concatenate([norm_s, norm_a, norm_tr], axis=-1)
372
+ node_feats = self.node_proj(self.out_norm(norms))
373
+
374
+ if isinstance(self.readout, WeightedAtomReadOut):
375
+ out = self.readout(node_feats, batch=batch, num_graphs=n_graphs)
376
+ else:
377
+ pooled = self.readout(node_feats, batch=batch, num_graphs=n_graphs)
378
+ out = self.final_mlp(pooled)
379
+
380
+ return ops.squeeze(out, axis=-1)
381
+
@@ -0,0 +1,167 @@
1
+ """Unit tests for multi-backend materials models."""
2
+
3
+ import pytest
4
+ import numpy as np
5
+ import keras
6
+ from keras import ops
7
+
8
+ from k3_node.applications.materials import (
9
+ MEGNet,
10
+ M3GNet,
11
+ TensorNet,
12
+ CHGNet,
13
+ SO3Net,
14
+ GRACE,
15
+ QET,
16
+ TransformedTargetModel,
17
+ Potential,
18
+ BondExpansion,
19
+ RadialBesselFunction,
20
+ FourierExpansion,
21
+ ChebyshevRadialBasis,
22
+ RealSphericalHarmonics,
23
+ LinearQeq,
24
+ get_available_pretrained_models,
25
+ )
26
+
27
+
28
+ @pytest.fixture
29
+ def synthetic_crystal():
30
+ """Create a synthetic crystal graph for testing."""
31
+ num_nodes = 4
32
+ pos = np.array([
33
+ [0.0, 0.0, 0.0],
34
+ [1.0, 0.5, 0.0],
35
+ [0.5, 1.2, 0.8],
36
+ [1.5, 1.5, 1.0],
37
+ ], dtype=np.float32)
38
+ edge_index = np.array([
39
+ [0, 1, 1, 2, 2, 3, 3, 0],
40
+ [1, 0, 2, 1, 3, 2, 0, 3],
41
+ ], dtype=np.int32)
42
+ line_edge_index = np.array([
43
+ [0, 1, 2, 3],
44
+ [1, 2, 3, 0],
45
+ ], dtype=np.int32)
46
+ node_type = np.array([6, 8, 1, 6], dtype=np.int32)
47
+ batch = np.array([0, 0, 0, 0], dtype=np.int32)
48
+ state_attr = np.array([[0.0, 0.0]], dtype=np.float32)
49
+ return {
50
+ "pos": pos,
51
+ "edge_index": edge_index,
52
+ "line_edge_index": line_edge_index,
53
+ "node_type": node_type,
54
+ "batch": batch,
55
+ "state_attr": state_attr,
56
+ }
57
+
58
+
59
+ def test_megnet_forward(synthetic_crystal):
60
+ model = MEGNet(
61
+ dim_node_embedding=16,
62
+ dim_edge_embedding=20,
63
+ dim_state_embedding=2,
64
+ nblocks=2,
65
+ hidden_layer_sizes_input=(32, 16),
66
+ hidden_layer_sizes_conv=(32, 16),
67
+ hidden_layer_sizes_output=(16,),
68
+ )
69
+ out = model(synthetic_crystal)
70
+ assert out.shape == () or out.shape == (1,)
71
+ assert np.isfinite(ops.convert_to_numpy(out)).all()
72
+
73
+
74
+ def test_m3gnet_forward(synthetic_crystal):
75
+ model = M3GNet(
76
+ dim_node_embedding=16,
77
+ dim_edge_embedding=16,
78
+ nblocks=2,
79
+ units=16,
80
+ max_n=3,
81
+ max_l=3,
82
+ )
83
+ out = model(synthetic_crystal)
84
+ assert out.shape == () or out.shape == (1,)
85
+ assert np.isfinite(ops.convert_to_numpy(out)).all()
86
+
87
+
88
+ def test_tensornet_forward(synthetic_crystal):
89
+ model = TensorNet(
90
+ units=16,
91
+ nblocks=2,
92
+ num_rbf=16,
93
+ )
94
+ out = model(synthetic_crystal)
95
+ assert out.shape == () or out.shape == (1,)
96
+ assert np.isfinite(ops.convert_to_numpy(out)).all()
97
+
98
+
99
+ def test_chgnet_forward(synthetic_crystal):
100
+ model = CHGNet(
101
+ dim_atom_embedding=16,
102
+ dim_bond_embedding=16,
103
+ dim_angle_embedding=16,
104
+ num_blocks=2,
105
+ atom_conv_hidden_dims=(16,),
106
+ bond_conv_hidden_dims=(16,),
107
+ )
108
+ out = model(synthetic_crystal)
109
+ assert out.shape == () or out.shape == (1,)
110
+ assert np.isfinite(ops.convert_to_numpy(out)).all()
111
+
112
+
113
+ def test_so3net_forward(synthetic_crystal):
114
+ model = SO3Net(
115
+ units=16,
116
+ nblocks=2,
117
+ lmax=2,
118
+ num_rbf=16,
119
+ )
120
+ out = model(synthetic_crystal)
121
+ assert out.shape == () or out.shape == (1,)
122
+ assert np.isfinite(ops.convert_to_numpy(out)).all()
123
+
124
+
125
+ def test_grace_forward(synthetic_crystal):
126
+ model = GRACE(
127
+ cutoff=5.0,
128
+ n_rad_base=6,
129
+ lmax=2,
130
+ embedding_size=8,
131
+ max_order=2,
132
+ nblocks=2,
133
+ readout_hidden=(16,),
134
+ )
135
+ out = model(synthetic_crystal)
136
+ assert out.shape == () or out.shape == (1,)
137
+ assert np.isfinite(ops.convert_to_numpy(out)).all()
138
+
139
+
140
+ def test_qet_forward(synthetic_crystal):
141
+ model = QET(
142
+ units=16,
143
+ nblocks=2,
144
+ num_rbf=16,
145
+ )
146
+ out = model(synthetic_crystal)
147
+ assert out.shape == () or out.shape == (1,)
148
+ assert np.isfinite(ops.convert_to_numpy(out)).all()
149
+
150
+
151
+ def test_wrappers(synthetic_crystal):
152
+ base_model = MEGNet(dim_node_embedding=8, dim_edge_embedding=16, nblocks=1)
153
+ tt_model = TransformedTargetModel(model=base_model, mean=5.0, std=2.0)
154
+ out_tt = tt_model(synthetic_crystal)
155
+ assert np.isfinite(ops.convert_to_numpy(out_tt)).all()
156
+
157
+ pot = Potential(model=base_model, data_mean=-1.5, data_std=0.8)
158
+ out_pot = pot(synthetic_crystal)
159
+ assert np.isfinite(ops.convert_to_numpy(out_pot)).all()
160
+
161
+
162
+ def test_available_models():
163
+ models = get_available_pretrained_models()
164
+ assert len(models) > 0
165
+ assert any("MEGNet" in m for m in models)
166
+ assert any("M3GNet" in m for m in models)
167
+
@@ -0,0 +1,95 @@
1
+ """Model wrappers: TransformedTargetModel and Potential interatomic potential."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Optional, Union, Dict, Any
6
+ import keras
7
+ from keras import layers, ops
8
+ import numpy as np
9
+
10
+
11
+ class TransformedTargetModel(keras.Model):
12
+ """Wraps a model and applies inverse transformation to predictions (e.g., mean/std denormalization).
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.models import MEGNet, TransformedTargetModel
18
+
19
+ # A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
20
+ structure = {
21
+ "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"),
22
+ "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
23
+ "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), # bond pairs forming angles
24
+ "node_type": np.array([6, 8, 1, 6]), # atomic numbers
25
+ "batch": np.zeros(4, dtype="int32"), # all atoms belong to structure 0
26
+ "state_attr": np.zeros((1, 2), dtype="float32"), # global state features
27
+ }
28
+
29
+ base = MEGNet(dim_node_embedding=8, dim_edge_embedding=16, nblocks=1)
30
+ model = TransformedTargetModel(model=base, mean=5.0, std=2.0) # outputs base * std + mean
31
+ print(tuple(model(structure).shape)) # (1,)
32
+ ```
33
+ """
34
+
35
+ def __init__(
36
+ self,
37
+ model: keras.Model,
38
+ mean: float = 0.0,
39
+ std: float = 1.0,
40
+ **kwargs,
41
+ ):
42
+ super().__init__(**kwargs)
43
+ self.model = model
44
+ self.mean = float(mean)
45
+ self.std = float(std)
46
+
47
+ def call(self, inputs, training=None, **kwargs):
48
+ pred = self.model(inputs, training=training, **kwargs)
49
+ return pred * self.std + self.mean
50
+
51
+
52
+ class Potential(keras.Model):
53
+ """Interatomic potential wrapping an energy model and computing energies, forces, and stresses.
54
+
55
+ Example:
56
+ ```python
57
+ import numpy as np
58
+ from k3_node.models import MEGNet, Potential
59
+
60
+ # A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
61
+ structure = {
62
+ "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"),
63
+ "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
64
+ "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), # bond pairs forming angles
65
+ "node_type": np.array([6, 8, 1, 6]), # atomic numbers
66
+ "batch": np.zeros(4, dtype="int32"), # all atoms belong to structure 0
67
+ "state_attr": np.zeros((1, 2), dtype="float32"), # global state features
68
+ }
69
+
70
+ base = MEGNet(dim_node_embedding=8, dim_edge_embedding=16, nblocks=1)
71
+ potential = Potential(model=base, data_mean=-1.5, data_std=0.8) # interatomic potential wrapper
72
+ print(tuple(potential(structure).shape)) # (1,)
73
+ ```
74
+ """
75
+
76
+ def __init__(
77
+ self,
78
+ model: keras.Model,
79
+ data_mean: float = 0.0,
80
+ data_std: float = 1.0,
81
+ element_refs: Optional[Dict[int, float]] = None,
82
+ calc_forces: bool = True,
83
+ **kwargs,
84
+ ):
85
+ super().__init__(**kwargs)
86
+ self.model = model
87
+ self.data_mean = float(data_mean)
88
+ self.data_std = float(data_std)
89
+ self.element_refs = element_refs or {}
90
+ self.calc_forces = calc_forces
91
+
92
+ def call(self, inputs, edge_index=None, node_type=None, training=None, **kwargs):
93
+ e_pred = self.model(inputs, edge_index=edge_index, node_type=node_type, training=training, **kwargs)
94
+ e_total = e_pred * self.data_std + self.data_mean
95
+ return e_total
@@ -0,0 +1,47 @@
1
+ from k3_node.data.batch import Batch, HeteroBatch
2
+ from k3_node.data.collate import collate
3
+ from k3_node.data.data import BaseData, Data
4
+ from k3_node.data.database import Database, RocksDatabase, SQLiteDatabase
5
+ from k3_node.data.dataset import Dataset
6
+ from k3_node.data.download import download_google_url, download_url
7
+ from k3_node.data.extract import extract_bz2, extract_gz, extract_tar, extract_zip
8
+ from k3_node.data.feature_store import FeatureStore, TensorAttr
9
+ from k3_node.data.graph_store import EdgeAttr, EdgeLayout, GraphStore
10
+ from k3_node.data.hetero_data import HeteroData
11
+ from k3_node.data.hypergraph_data import HyperGraphData, HypergraphData
12
+ from k3_node.data.in_memory_dataset import InMemoryDataset
13
+ from k3_node.data.makedirs import makedirs
14
+ from k3_node.data.on_disk_dataset import OnDiskDataset
15
+ from k3_node.data.separate import separate
16
+ from k3_node.data.temporal import TemporalData
17
+
18
+ __all__ = [
19
+ "Data",
20
+ "HeteroData",
21
+ "Batch",
22
+ "HeteroBatch",
23
+ "TemporalData",
24
+ "HypergraphData",
25
+ "HyperGraphData",
26
+ "Dataset",
27
+ "InMemoryDataset",
28
+ "OnDiskDataset",
29
+ "FeatureStore",
30
+ "GraphStore",
31
+ "TensorAttr",
32
+ "EdgeAttr",
33
+ "EdgeLayout",
34
+ "Database",
35
+ "SQLiteDatabase",
36
+ "RocksDatabase",
37
+ "makedirs",
38
+ "download_url",
39
+ "download_google_url",
40
+ "extract_tar",
41
+ "extract_zip",
42
+ "extract_bz2",
43
+ "extract_gz",
44
+ "collate",
45
+ "separate",
46
+ ]
47
+