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,40 @@
1
+ from keras import ops
2
+
3
+ from k3_node.models import CorrectAndSmooth
4
+
5
+
6
+ def test_correct_and_smooth():
7
+ y_soft = ops.repeat(ops.convert_to_tensor([[0.1, 0.5, 0.4]]), 6, axis=0)
8
+ y_true = ops.convert_to_tensor([1, 0, 0, 2, 1, 1])
9
+ edge_index = ops.convert_to_tensor([[0, 1, 1, 2, 4, 5], [1, 0, 2, 1, 5, 4]])
10
+ mask = ops.convert_to_tensor([True, False, True, False, True, False])
11
+
12
+ model = CorrectAndSmooth(
13
+ num_correction_layers=2,
14
+ correction_alpha=0.5,
15
+ num_smoothing_layers=2,
16
+ smoothing_alpha=0.5,
17
+ )
18
+ assert str(model) == ('CorrectAndSmooth(\n'
19
+ ' correct: num_layers=2, alpha=0.5\n'
20
+ ' smooth: num_layers=2, alpha=0.5\n'
21
+ ' autoscale=True, scale=1.0\n'
22
+ ')')
23
+
24
+ out = model.correct(y_soft, y_true[mask], mask, edge_index)
25
+ assert ops.shape(out) == (6, 3)
26
+
27
+ out = model.smooth(y_soft, y_true[mask], mask, edge_index)
28
+ assert ops.shape(out) == (6, 3)
29
+
30
+ # Without autoscale:
31
+ model_no_auto = CorrectAndSmooth(
32
+ num_correction_layers=2,
33
+ correction_alpha=0.5,
34
+ num_smoothing_layers=2,
35
+ smoothing_alpha=0.5,
36
+ autoscale=False,
37
+ )
38
+ out = model_no_auto.correct(y_soft, y_true[mask], mask, edge_index)
39
+ assert ops.shape(out) == (6, 3)
40
+
@@ -0,0 +1,68 @@
1
+ import numpy as np
2
+ from keras import ops
3
+
4
+ from k3_node.models import DeepGraphInfomax
5
+
6
+
7
+ def test_infomax():
8
+ model = DeepGraphInfomax(
9
+ hidden_channels=16,
10
+ encoder=lambda x: x,
11
+ summary=lambda z, *args: ops.mean(z, axis=0),
12
+ corruption=lambda x: x + 1,
13
+ )
14
+ assert str(model) == "DeepGraphInfomax(16)"
15
+
16
+ x = ops.ones((20, 16))
17
+
18
+ pos_z, neg_z, summary = model(x)
19
+ assert ops.shape(pos_z) == (20, 16)
20
+ assert ops.shape(neg_z) == (20, 16)
21
+ assert ops.shape(summary) == (16,)
22
+
23
+ loss = model.loss(pos_z, neg_z, summary)
24
+ assert float(ops.convert_to_numpy(loss)) >= 0
25
+
26
+ acc = model.test(
27
+ train_z=ops.ones((20, 16)),
28
+ train_y=ops.convert_to_tensor(np.random.randint(0, 10, (20,))),
29
+ test_z=ops.ones((20, 16)),
30
+ test_y=ops.convert_to_tensor(np.random.randint(0, 10, (20,))),
31
+ )
32
+ assert 0 <= acc <= 1
33
+
34
+
35
+ def test_infomax_predefined_model():
36
+ from k3_node.layers.conv import GCNConv
37
+
38
+ class Encoder:
39
+ def __init__(self):
40
+ self.conv1 = GCNConv(16, 16)
41
+ self.conv2 = GCNConv(16, 16)
42
+
43
+ def __call__(self, x, edge_index, edge_weight=None):
44
+ x = ops.relu(self.conv1(x, edge_index, edge_weight=edge_weight))
45
+ return self.conv2(x, edge_index, edge_weight=edge_weight)
46
+
47
+ def corruption(x, edge_index, edge_weight):
48
+ perm = np.random.permutation(ops.shape(x)[0])
49
+ return ops.take(x, ops.convert_to_tensor(perm), axis=0), edge_index, edge_weight
50
+
51
+ model = DeepGraphInfomax(
52
+ hidden_channels=16,
53
+ encoder=Encoder(),
54
+ summary=lambda z, *args, **kwargs: ops.sigmoid(ops.mean(z, axis=0)),
55
+ corruption=corruption,
56
+ )
57
+
58
+ x = ops.convert_to_tensor(np.random.randn(4, 16).astype("float32"))
59
+ edge_index = ops.convert_to_tensor([[0, 0, 0, 1, 2, 3], [1, 2, 3, 0, 0, 0]], dtype="int64")
60
+ edge_weight = ops.convert_to_tensor(np.random.rand(edge_index.shape[1]).astype("float32"))
61
+
62
+ pos_z, neg_z, summary = model(x, edge_index, edge_weight=edge_weight)
63
+ assert ops.shape(pos_z) == (4, 16)
64
+ assert ops.shape(neg_z) == (4, 16)
65
+ assert ops.shape(summary) == (16,)
66
+
67
+ loss = model.loss(pos_z, neg_z, summary)
68
+ assert float(ops.convert_to_numpy(loss)) >= 0
@@ -0,0 +1,21 @@
1
+ import pytest
2
+ from keras import ops
3
+
4
+ from k3_node.layers.conv import GENConv
5
+ from k3_node.layers.norm import LayerNorm
6
+ from k3_node.models import DeepGCNLayer
7
+
8
+
9
+ @pytest.mark.parametrize("block_tuple", [("res+", 1), ("res", 1), ("dense", 2), ("plain", 1)])
10
+ def test_deepgcn(block_tuple):
11
+ block, expansion = block_tuple
12
+ x = ops.convert_to_tensor([[1.0] * 8] * 3, dtype="float32")
13
+ edge_index = ops.convert_to_tensor([[0, 1, 1, 2], [1, 0, 2, 1]], dtype="int64")
14
+ conv = GENConv(8, 8)
15
+ norm = LayerNorm(8)
16
+ act = ops.relu
17
+ layer = DeepGCNLayer(conv, norm, act, block=block)
18
+ assert str(layer) == f"DeepGCNLayer(block={block})"
19
+
20
+ out = layer(x, edge_index)
21
+ assert ops.shape(out) == (3, 8 * expansion)
@@ -0,0 +1,86 @@
1
+ import keras.ops as ops
2
+ from k3_node.models.dimenet import (
3
+ DimeNet,
4
+ DimeNetPlusPlus,
5
+ BesselBasisLayer,
6
+ SphericalBasisLayer,
7
+ triplets,
8
+ )
9
+
10
+
11
+ def test_triplets():
12
+ edge_index = ops.convert_to_tensor([
13
+ [0, 1, 1, 2, 0, 2],
14
+ [1, 0, 2, 1, 2, 0],
15
+ ])
16
+ col, row, idx_i, idx_j, idx_k, idx_kj, idx_ji = triplets(edge_index, num_nodes=3)
17
+ assert len(idx_i) > 0
18
+ assert len(idx_kj) == len(idx_ji)
19
+
20
+
21
+ def test_bessel_basis_layer():
22
+ bessel = BesselBasisLayer(num_radial=6, cutoff=5.0)
23
+ dist = ops.convert_to_tensor([1.0, 2.0, 3.0])
24
+ out = bessel(dist)
25
+ assert out.shape == (3, 6)
26
+
27
+
28
+ def test_spherical_basis_layer():
29
+ sbf = SphericalBasisLayer(num_spherical=3, num_radial=6, cutoff=5.0)
30
+ dist = ops.convert_to_tensor([1.0, 2.0])
31
+ angle = ops.convert_to_tensor([0.5, 1.2])
32
+ idx_kj = ops.convert_to_tensor([0, 1])
33
+ out = sbf(dist, angle, idx_kj)
34
+ assert out.shape == (2, 18)
35
+
36
+
37
+ def test_dimenet():
38
+ z = ops.convert_to_tensor([1, 6, 8, 1])
39
+ pos = ops.convert_to_tensor([
40
+ [0.0, 0.0, 0.0],
41
+ [1.0, 0.0, 0.0],
42
+ [0.0, 1.0, 0.0],
43
+ [1.0, 1.0, 0.0],
44
+ ], dtype="float32")
45
+
46
+ model = DimeNet(
47
+ hidden_channels=16,
48
+ out_channels=1,
49
+ num_blocks=2,
50
+ num_bilinear=8,
51
+ num_spherical=3,
52
+ num_radial=6,
53
+ cutoff=5.0,
54
+ )
55
+ out = model(z, pos)
56
+ assert out.shape == (1,)
57
+
58
+ # With batch
59
+ batch = ops.convert_to_tensor([0, 0, 1, 1])
60
+ out_b = model(z, pos, batch=batch)
61
+ assert out_b.shape == (2, 1)
62
+
63
+
64
+ def test_dimenet_plus_plus():
65
+ z = ops.convert_to_tensor([1, 6, 8, 1])
66
+ pos = ops.convert_to_tensor([
67
+ [0.0, 0.0, 0.0],
68
+ [1.0, 0.0, 0.0],
69
+ [0.0, 1.0, 0.0],
70
+ [1.0, 1.0, 0.0],
71
+ ], dtype="float32")
72
+
73
+ model = DimeNetPlusPlus(
74
+ hidden_channels=16,
75
+ out_channels=1,
76
+ num_blocks=2,
77
+ int_emb_size=8,
78
+ basis_emb_size=8,
79
+ out_emb_channels=16,
80
+ num_spherical=3,
81
+ num_radial=6,
82
+ cutoff=5.0,
83
+ )
84
+ out = model(z, pos)
85
+ assert out.shape == (1,)
86
+
@@ -0,0 +1,138 @@
1
+ """Test domain-specific API organization (applications: materials, bio, chemistry)."""
2
+
3
+ import pytest
4
+
5
+
6
+ def test_clean_main_api():
7
+ """Ensure main k3_node API is clean and domain packages are under applications."""
8
+ import k3_node
9
+
10
+ assert not hasattr(k3_node, "materials"), "k3_node should not expose materials directly"
11
+ assert not hasattr(k3_node, "bio"), "k3_node should not expose bio directly"
12
+ assert not hasattr(k3_node, "chemistry"), "k3_node should not expose chemistry directly"
13
+
14
+ assert hasattr(k3_node, "applications"), "k3_node must expose applications"
15
+ assert hasattr(k3_node.applications, "materials"), "applications must expose materials"
16
+ assert hasattr(k3_node.applications, "bio"), "applications must expose bio"
17
+ assert hasattr(k3_node.applications, "chemistry"), "applications must expose chemistry"
18
+
19
+
20
+ def test_materials_application_api():
21
+ # Via k3_node.applications.materials
22
+ from k3_node.applications.materials import (
23
+ MEGNet,
24
+ M3GNet,
25
+ TensorNet,
26
+ CHGNet,
27
+ SO3Net,
28
+ GRACE,
29
+ QET,
30
+ TransformedTargetModel,
31
+ Potential,
32
+ download_matgl_checkpoint,
33
+ load_matgl_weights,
34
+ load_model,
35
+ get_available_pretrained_models,
36
+ basis,
37
+ core,
38
+ readout,
39
+ )
40
+ assert MEGNet is not None
41
+ assert M3GNet is not None
42
+ assert TensorNet is not None
43
+ assert CHGNet is not None
44
+ assert SO3Net is not None
45
+ assert GRACE is not None
46
+ assert QET is not None
47
+ assert basis is not None
48
+ assert core is not None
49
+ assert readout is not None
50
+
51
+ # Via k3_node.models.materials (backward compatibility)
52
+ from k3_node.models.materials import (
53
+ MEGNet as MatMEGNet,
54
+ M3GNet as MatM3GNet,
55
+ TensorNet as MatTensorNet,
56
+ )
57
+ assert MatMEGNet is MEGNet
58
+ assert MatM3GNet is M3GNet
59
+ assert MatTensorNet is TensorNet
60
+
61
+
62
+ def test_bio_application_api():
63
+ # Via k3_node.applications.bio
64
+ from k3_node.applications.bio import (
65
+ UniMolDockingModel,
66
+ DockingPoseModelV2,
67
+ download_unimol_checkpoint,
68
+ load_unimol_weights,
69
+ download_unimol_docking_checkpoint,
70
+ load_unimol_docking_weights,
71
+ )
72
+ assert UniMolDockingModel is not None
73
+ assert DockingPoseModelV2 is not None
74
+
75
+ # Via k3_node.models.bio (backward compatibility)
76
+ from k3_node.models.bio import (
77
+ UniMolDockingModel as BioDocking1,
78
+ DockingPoseModelV2 as BioDocking2,
79
+ )
80
+ assert BioDocking1 is UniMolDockingModel
81
+ assert BioDocking2 is DockingPoseModelV2
82
+
83
+
84
+ def test_chemistry_application_api():
85
+ # Via k3_node.applications.chemistry
86
+ from k3_node.applications.chemistry import (
87
+ AttentiveFP,
88
+ DimeNet,
89
+ DimeNetPlusPlus,
90
+ GROVER,
91
+ MoleBERT,
92
+ NeuralFingerprint,
93
+ SchNet,
94
+ ViSNet,
95
+ GNNFF,
96
+ Graphormer,
97
+ Graphormer3D,
98
+ UniMolModel,
99
+ UniMolConfGenModel,
100
+ UniMol2Model,
101
+ UniMolPlusPCQModel,
102
+ UniMolPlusOC20Model,
103
+ )
104
+ assert AttentiveFP is not None
105
+ assert DimeNet is not None
106
+ assert SchNet is not None
107
+ assert UniMolModel is not None
108
+
109
+ # Via k3_node.models.chemistry (backward compatibility)
110
+ from k3_node.models.chemistry import (
111
+ SchNet as ChemSchNet,
112
+ DimeNet as ChemDimeNet,
113
+ UniMolModel as ChemUniMol,
114
+ )
115
+ assert ChemSchNet is SchNet
116
+ assert ChemDimeNet is DimeNet
117
+ assert ChemUniMol is UniMolModel
118
+
119
+
120
+ def test_backward_compatibility():
121
+ # All models still exportable from top-level k3_node.models
122
+ from k3_node.models import (
123
+ MEGNet,
124
+ M3GNet,
125
+ TensorNet,
126
+ CHGNet,
127
+ SO3Net,
128
+ GRACE,
129
+ QET,
130
+ SchNet,
131
+ DimeNet,
132
+ UniMolModel,
133
+ UniMolDockingModel,
134
+ DockingPoseModelV2,
135
+ )
136
+ assert MEGNet is not None
137
+ assert SchNet is not None
138
+ assert UniMolModel is not None
@@ -0,0 +1,24 @@
1
+ import keras.ops as ops
2
+ from k3_node.models.gnnff import GNNFF, GaussianFilter
3
+
4
+
5
+ def test_gaussian_filter():
6
+ gf = GaussianFilter(start=0.0, stop=5.0, num_gaussians=10)
7
+ dist = ops.convert_to_tensor([0.5, 1.5, 3.0])
8
+ out = gf(dist)
9
+ assert out.shape == (3, 10)
10
+
11
+
12
+ def test_gnnff():
13
+ z = ops.convert_to_tensor([1, 6, 8, 1])
14
+ pos = ops.convert_to_tensor([
15
+ [0.0, 0.0, 0.0],
16
+ [1.0, 0.0, 0.0],
17
+ [0.0, 1.0, 0.0],
18
+ [1.0, 1.0, 0.0],
19
+ ], dtype="float32")
20
+
21
+ model = GNNFF(hidden_node_channels=16, hidden_edge_channels=16, num_layers=2)
22
+ force = model(z, pos)
23
+ assert force.shape == (4, 3)
24
+
@@ -0,0 +1,271 @@
1
+ import os
2
+ import tempfile
3
+ import numpy as np
4
+ import pytest
5
+ import keras
6
+ from keras import ops
7
+
8
+ from k3_node.models.gps_model import (
9
+ AtomEncoder,
10
+ BondEncoder,
11
+ RWSEEncoder,
12
+ CustomGatedGCN,
13
+ GPSLayer,
14
+ SANGraphHead,
15
+ GPSModel,
16
+ load_gps_weights,
17
+ download_gps_checkpoint,
18
+ )
19
+
20
+
21
+ def test_atom_and_bond_encoders():
22
+ # 9 features for atoms, 3 for bonds
23
+ x = ops.convert_to_tensor(np.array([[6, 0, 4, 4, 3, 2, 2, 0, 0],
24
+ [8, 0, 2, 2, 2, 1, 1, 0, 0],
25
+ [1, 0, 0, 0, 0, 0, 0, 0, 0]], dtype=np.int32))
26
+ atom_enc = AtomEncoder(emb_dim=64)
27
+ x_emb = atom_enc(x)
28
+ assert ops.shape(x_emb) == (3, 64)
29
+
30
+ edge_attr = ops.convert_to_tensor(np.array([[0, 0, 0],
31
+ [1, 0, 0]], dtype=np.int32))
32
+ bond_enc = BondEncoder(emb_dim=64)
33
+ e_emb = bond_enc(edge_attr)
34
+ assert ops.shape(e_emb) == (2, 64)
35
+
36
+
37
+ def test_rwse_encoder():
38
+ pestat = ops.convert_to_tensor(np.ones((4, 16), dtype=np.float32))
39
+ rwse = RWSEEncoder(num_rw_steps=16, pe_dim=20)
40
+ pe_out = rwse(pestat, training=False)
41
+ assert ops.shape(pe_out) == (4, 20)
42
+
43
+
44
+ def test_custom_gated_gcn():
45
+ x = ops.convert_to_tensor(np.random.randn(4, 32).astype(np.float32))
46
+ e = ops.convert_to_tensor(np.random.randn(6, 32).astype(np.float32))
47
+ edge_index = ops.convert_to_tensor(np.array([[0, 1, 1, 2, 2, 3],
48
+ [1, 0, 2, 1, 3, 2]], dtype=np.int32))
49
+
50
+ layer = CustomGatedGCN(in_dim=32, out_dim=32, dropout=0.0, residual=True, act="gelu")
51
+ x_out, e_out = layer(x, edge_index, e, training=False)
52
+
53
+ assert ops.shape(x_out) == (4, 32)
54
+ assert ops.shape(e_out) == (6, 32)
55
+
56
+
57
+ def test_gps_layer():
58
+ dim_h = 32
59
+ num_heads = 4
60
+ x = ops.convert_to_tensor(np.random.randn(5, dim_h).astype(np.float32))
61
+ e = ops.convert_to_tensor(np.random.randn(6, dim_h).astype(np.float32))
62
+ edge_index = ops.convert_to_tensor(np.array([[0, 1, 1, 2, 3, 4],
63
+ [1, 0, 2, 1, 4, 3]], dtype=np.int32))
64
+ # Two graphs: graph 0 has 3 nodes, graph 1 has 2 nodes
65
+ batch = ops.convert_to_tensor(np.array([0, 0, 0, 1, 1], dtype=np.int32))
66
+
67
+ layer = GPSLayer(
68
+ dim_h=dim_h,
69
+ local_gnn_type="CustomGatedGCN",
70
+ global_model_type="Transformer",
71
+ num_heads=num_heads,
72
+ act="gelu",
73
+ dropout=0.0,
74
+ attn_dropout=0.0,
75
+ batch_norm=True,
76
+ )
77
+
78
+ x_out, e_out = layer(x, edge_index, e, batch=batch, training=False)
79
+ assert ops.shape(x_out) == (5, dim_h)
80
+ assert ops.shape(e_out) == (6, dim_h)
81
+
82
+
83
+ def test_san_graph_head():
84
+ x = ops.convert_to_tensor(np.random.randn(6, 64).astype(np.float32))
85
+ batch = ops.convert_to_tensor(np.array([0, 0, 0, 1, 1, 1], dtype=np.int32))
86
+
87
+ head = SANGraphHead(dim_in=64, dim_out=1, L=2, act="gelu", pooling="mean")
88
+ pred = head(x, batch=batch, training=False)
89
+ assert ops.shape(pred) == (2, 1)
90
+
91
+ head_sum = SANGraphHead(dim_in=64, dim_out=2, L=1, act="relu", pooling="sum")
92
+ pred_sum = head_sum(x, batch=batch, training=False)
93
+ assert ops.shape(pred_sum) == (2, 2)
94
+
95
+
96
+ def test_gps_model_forward():
97
+ model = GPSModel(
98
+ dim_in=32,
99
+ dim_out=1,
100
+ num_layers=2,
101
+ dim_hidden=32,
102
+ num_heads=4,
103
+ local_gnn_type="CustomGatedGCN",
104
+ act="gelu",
105
+ dropout=0.1,
106
+ attn_dropout=0.1,
107
+ batch_norm=True,
108
+ node_encoder_type="Atom+RWSE",
109
+ edge_encoder_type="Bond",
110
+ rwse_num_steps=16,
111
+ rwse_dim_pe=8,
112
+ graph_pooling="mean",
113
+ head_layers=2,
114
+ )
115
+
116
+ x = ops.convert_to_tensor(np.array([[6, 0, 4, 4, 3, 2, 2, 0, 0],
117
+ [8, 0, 2, 2, 2, 1, 1, 0, 0],
118
+ [6, 0, 3, 3, 3, 2, 2, 0, 0]], dtype=np.int32))
119
+ edge_index = ops.convert_to_tensor(np.array([[0, 1, 1, 2],
120
+ [1, 0, 2, 1]], dtype=np.int32))
121
+ edge_attr = ops.convert_to_tensor(np.array([[0, 0, 0],
122
+ [0, 0, 0],
123
+ [1, 0, 0],
124
+ [1, 0, 0]], dtype=np.int32))
125
+ pestat = ops.convert_to_tensor(np.ones((3, 16), dtype=np.float32))
126
+ batch = ops.convert_to_tensor(np.array([0, 0, 0], dtype=np.int32))
127
+
128
+ pred_eval = model(x, edge_index, edge_attr=edge_attr, pestat_RWSE=pestat, batch=batch, training=False)
129
+ pred_train = model(x, edge_index, edge_attr=edge_attr, pestat_RWSE=pestat, batch=batch, training=True)
130
+
131
+ assert ops.shape(pred_eval) == (1, 1)
132
+ assert ops.shape(pred_train) == (1, 1)
133
+
134
+
135
+ def test_gps_model_load_weights_synthetic():
136
+ torch = pytest.importorskip("torch")
137
+
138
+ dim_h = 32
139
+ rwse_pe = 8
140
+ model = GPSModel(
141
+ dim_in=dim_h,
142
+ dim_out=1,
143
+ num_layers=1,
144
+ dim_hidden=dim_h,
145
+ num_heads=2,
146
+ local_gnn_type="CustomGatedGCN",
147
+ act="gelu",
148
+ dropout=0.0,
149
+ attn_dropout=0.0,
150
+ batch_norm=True,
151
+ node_encoder_type="Atom+RWSE",
152
+ edge_encoder_type="Bond",
153
+ rwse_num_steps=16,
154
+ rwse_dim_pe=rwse_pe,
155
+ graph_pooling="mean",
156
+ head_layers=1,
157
+ )
158
+ model.build(None)
159
+
160
+ # Synthetic PyTorch state dict
161
+ state_dict = {}
162
+ for i in range(9):
163
+ dim_val = [119, 4, 12, 12, 10, 6, 6, 2, 2][i]
164
+ state_dict[f"encoder.node_encoder.encoder1.atom_embedding_list.{i}.weight"] = torch.randn(dim_val, dim_h - rwse_pe)
165
+ for i in range(3):
166
+ dim_val = [5, 6, 2][i]
167
+ state_dict[f"encoder.edge_encoder.bond_embedding_list.{i}.weight"] = torch.randn(dim_val, dim_h)
168
+
169
+ p_rwse = "encoder.node_encoder.encoder2"
170
+ state_dict[f"{p_rwse}.raw_norm.weight"] = torch.ones(16)
171
+ state_dict[f"{p_rwse}.raw_norm.bias"] = torch.zeros(16)
172
+ state_dict[f"{p_rwse}.raw_norm.running_mean"] = torch.zeros(16)
173
+ state_dict[f"{p_rwse}.raw_norm.running_var"] = torch.ones(16)
174
+ state_dict[f"{p_rwse}.pe_encoder.weight"] = torch.randn(rwse_pe, 16)
175
+ state_dict[f"{p_rwse}.pe_encoder.bias"] = torch.zeros(rwse_pe)
176
+
177
+ p0 = "layers.0"
178
+ for proj in ["A", "B", "C", "D", "E"]:
179
+ state_dict[f"{p0}.local_model.{proj}.weight"] = torch.randn(dim_h, dim_h)
180
+ state_dict[f"{p0}.local_model.{proj}.bias"] = torch.zeros(dim_h)
181
+ for bn in ["bn_node_x", "bn_edge_e"]:
182
+ state_dict[f"{p0}.local_model.{bn}.weight"] = torch.ones(dim_h)
183
+ state_dict[f"{p0}.local_model.{bn}.bias"] = torch.zeros(dim_h)
184
+ state_dict[f"{p0}.local_model.{bn}.running_mean"] = torch.zeros(dim_h)
185
+ state_dict[f"{p0}.local_model.{bn}.running_var"] = torch.ones(dim_h)
186
+
187
+ state_dict[f"{p0}.norm1_local.weight"] = torch.ones(dim_h)
188
+ state_dict[f"{p0}.norm1_local.bias"] = torch.zeros(dim_h)
189
+ state_dict[f"{p0}.norm1_local.running_mean"] = torch.zeros(dim_h)
190
+ state_dict[f"{p0}.norm1_local.running_var"] = torch.ones(dim_h)
191
+
192
+ state_dict[f"{p0}.self_attn.in_proj_weight"] = torch.randn(3 * dim_h, dim_h)
193
+ state_dict[f"{p0}.self_attn.in_proj_bias"] = torch.zeros(3 * dim_h)
194
+ state_dict[f"{p0}.self_attn.out_proj.weight"] = torch.randn(dim_h, dim_h)
195
+ state_dict[f"{p0}.self_attn.out_proj.bias"] = torch.zeros(dim_h)
196
+
197
+ state_dict[f"{p0}.norm1_attn.weight"] = torch.ones(dim_h)
198
+ state_dict[f"{p0}.norm1_attn.bias"] = torch.zeros(dim_h)
199
+ state_dict[f"{p0}.norm1_attn.running_mean"] = torch.zeros(dim_h)
200
+ state_dict[f"{p0}.norm1_attn.running_var"] = torch.ones(dim_h)
201
+
202
+ state_dict[f"{p0}.ff_linear1.weight"] = torch.randn(dim_h * 2, dim_h)
203
+ state_dict[f"{p0}.ff_linear1.bias"] = torch.zeros(dim_h * 2)
204
+ state_dict[f"{p0}.ff_linear2.weight"] = torch.randn(dim_h, dim_h * 2)
205
+ state_dict[f"{p0}.ff_linear2.bias"] = torch.zeros(dim_h)
206
+
207
+ state_dict[f"{p0}.norm2.weight"] = torch.ones(dim_h)
208
+ state_dict[f"{p0}.norm2.bias"] = torch.zeros(dim_h)
209
+ state_dict[f"{p0}.norm2.running_mean"] = torch.zeros(dim_h)
210
+ state_dict[f"{p0}.norm2.running_var"] = torch.ones(dim_h)
211
+
212
+ state_dict["post_mp.FC_layers.0.weight"] = torch.randn(dim_h // 2, dim_h)
213
+ state_dict["post_mp.FC_layers.0.bias"] = torch.zeros(dim_h // 2)
214
+ state_dict["post_mp.FC_layers.1.weight"] = torch.randn(1, dim_h // 2)
215
+ state_dict["post_mp.FC_layers.1.bias"] = torch.zeros(1)
216
+
217
+ with tempfile.NamedTemporaryFile(suffix=".ckpt", delete=False) as f:
218
+ torch.save({"model_state": state_dict}, f.name)
219
+ ckpt_path = f.name
220
+
221
+ try:
222
+ load_gps_weights(model, ckpt_path)
223
+ finally:
224
+ if os.path.exists(ckpt_path):
225
+ os.remove(ckpt_path)
226
+
227
+
228
+ def test_gps_model_pretrained_checkpoint_if_available():
229
+ torch = pytest.importorskip("torch")
230
+ ckpt_path = download_gps_checkpoint("pcqm4m-GPS+RWSE.deep")
231
+ if os.path.exists(ckpt_path):
232
+ model = GPSModel(
233
+ dim_in=256,
234
+ dim_out=1,
235
+ num_layers=16,
236
+ dim_hidden=256,
237
+ num_heads=8,
238
+ local_gnn_type="CustomGatedGCN",
239
+ act="gelu",
240
+ dropout=0.0,
241
+ attn_dropout=0.0,
242
+ batch_norm=True,
243
+ layer_norm=False,
244
+ node_encoder_type="Atom+RWSE",
245
+ edge_encoder_type="Bond",
246
+ rwse_num_steps=16,
247
+ rwse_dim_pe=20,
248
+ graph_pooling="mean",
249
+ head_layers=2,
250
+ )
251
+ x = ops.convert_to_tensor(np.array([[6, 0, 4, 4, 3, 2, 2, 0, 0],
252
+ [8, 0, 2, 2, 2, 1, 1, 0, 0],
253
+ [6, 0, 3, 3, 3, 2, 2, 0, 0]], dtype=np.int32))
254
+ edge_index = ops.convert_to_tensor(np.array([[0, 1, 1, 2],
255
+ [1, 0, 2, 1]], dtype=np.int32))
256
+ edge_attr = ops.convert_to_tensor(np.array([[0, 0, 0],
257
+ [0, 0, 0],
258
+ [1, 0, 0],
259
+ [1, 0, 0]], dtype=np.int32))
260
+ pestat = ops.convert_to_tensor(np.ones((3, 16), dtype=np.float32))
261
+ batch = ops.convert_to_tensor(np.array([0, 0, 0], dtype=np.int32))
262
+
263
+ # Build and load
264
+ _ = model(x, edge_index, edge_attr=edge_attr, pestat_RWSE=pestat, batch=batch, training=False)
265
+ load_gps_weights(model, ckpt_path)
266
+
267
+ out = model(x, edge_index, edge_attr=edge_attr, pestat_RWSE=pestat, batch=batch, training=False)
268
+ val = float(ops.convert_to_numpy(out)[0, 0])
269
+ # PyTorch reference output was ~15.3017
270
+ np.testing.assert_allclose(val, 15.3017, rtol=1e-3, atol=1e-3)
271
+
@@ -0,0 +1,34 @@
1
+ import keras.ops as ops
2
+ from k3_node.models.gpse import GPSE, GPSENodeEncoder
3
+
4
+
5
+ def test_gpse_encoder():
6
+ x = ops.ones((4, 16))
7
+ pos_enc = ops.ones((4, 32))
8
+
9
+ encoder = GPSENodeEncoder(
10
+ dim_emb=64,
11
+ dim_pe_in=32,
12
+ dim_pe_out=16,
13
+ dim_in=16,
14
+ expand_x=True,
15
+ )
16
+ out = encoder(x, pos_enc)
17
+ assert out.shape == (4, 64)
18
+
19
+
20
+ def test_gpse_model():
21
+ x = ops.ones((4, 16))
22
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]])
23
+
24
+ model = GPSE(
25
+ dim_in=16,
26
+ dim_inner=32,
27
+ layers_pre_mp=1,
28
+ layers_mp=2,
29
+ layers_post_mp=1,
30
+ use_repr=True,
31
+ )
32
+ out = model(x, edge_index)
33
+ assert out.shape == (4, 32)
34
+