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,868 @@
1
+ import math
2
+ import os
3
+ from typing import Optional, Union, Tuple, List, Dict, Any
4
+
5
+ import keras
6
+ from keras import layers, ops
7
+
8
+
9
+ class GaussianLayer(layers.Layer):
10
+ r"""Gaussian basis function expansion over pairwise distances modulated by edge types.
11
+
12
+ Args:
13
+ num_kernel (int, optional): Number of Gaussian basis kernels. (default: ``128``)
14
+ edge_types (int, optional): Number of pairwise edge types. (default: ``4096``)
15
+ **kwargs: Additional layer arguments.
16
+
17
+ Example:
18
+ ```python
19
+ import numpy as np
20
+ from k3_node.models import GaussianLayer
21
+
22
+ dist = np.random.rand(2, 4, 4).astype("float32") * 3.0 # pairwise distances of 2 structures
23
+ edge_types = np.random.randint(0, 16, size=(2, 4, 4)) # type of each atom pair
24
+ layer = GaussianLayer(num_kernel=32, edge_types=16)
25
+ print(tuple(layer(dist, edge_types).shape)) # (2, 4, 4, 32): distance expansion per atom pair
26
+ ```
27
+ """
28
+
29
+ def __init__(self, num_kernel: int = 128, edge_types: int = 4096, **kwargs):
30
+ super().__init__(**kwargs)
31
+ self.num_kernel = num_kernel
32
+ self.edge_types = edge_types
33
+
34
+ self.means = layers.Embedding(
35
+ input_dim=1,
36
+ output_dim=num_kernel,
37
+ name="means",
38
+ )
39
+ self.stds = layers.Embedding(
40
+ input_dim=1,
41
+ output_dim=num_kernel,
42
+ name="stds",
43
+ )
44
+ self.mul = layers.Embedding(
45
+ input_dim=edge_types,
46
+ output_dim=1,
47
+ name="mul",
48
+ )
49
+ self.bias = layers.Embedding(
50
+ input_dim=edge_types,
51
+ output_dim=1,
52
+ name="bias",
53
+ )
54
+
55
+ def build(self, input_shape=None):
56
+ if not self.built:
57
+ self.means.build(None)
58
+ self.stds.build(None)
59
+ self.mul.build(None)
60
+ self.bias.build(None)
61
+ super().build(input_shape)
62
+
63
+ def call(self, dist, edge_types):
64
+ r"""
65
+ Args:
66
+ dist (Tensor): Pairwise distance matrix of shape ``[batch_size, num_nodes, num_nodes]``.
67
+ edge_types (Tensor): Pairwise edge type indices of shape ``[batch_size, num_nodes, num_nodes]``.
68
+
69
+ Returns:
70
+ Tensor: Gaussian basis expansion of shape ``[batch_size, num_nodes, num_nodes, num_kernel]``.
71
+ """
72
+ mul = self.mul(edge_types) # [B, N, N, 1]
73
+ bias = self.bias(edge_types) # [B, N, N, 1]
74
+ x = mul * ops.expand_dims(dist, axis=-1) + bias # [B, N, N, 1]
75
+
76
+ means = ops.reshape(self.means(ops.zeros((1,), dtype="int32")), (-1,)) # [K]
77
+ stds = ops.abs(ops.reshape(self.stds(ops.zeros((1,), dtype="int32")), (-1,))) + 1e-5 # [K]
78
+
79
+ a = math.sqrt(2 * math.pi)
80
+ diff = (x - means) / stds
81
+ return ops.exp(-0.5 * ops.power(diff, 2)) / (a * stds)
82
+
83
+
84
+ class RBF(layers.Layer):
85
+ r"""Radial Basis Function expansion over pairwise distances modulated by edge types.
86
+
87
+ Args:
88
+ num_kernel (int): Number of radial basis kernels.
89
+ edge_types (int): Number of edge types.
90
+ **kwargs: Additional layer arguments.
91
+
92
+ Example:
93
+ ```python
94
+ import numpy as np
95
+ from k3_node.models import RBF
96
+
97
+ dist = np.random.rand(2, 4, 4).astype("float32") * 3.0 # pairwise distances of 2 structures
98
+ edge_types = np.random.randint(0, 16, size=(2, 4, 4)) # type of each atom pair
99
+ layer = RBF(num_kernel=32, edge_types=16)
100
+ print(tuple(layer(dist, edge_types).shape)) # (2, 4, 4, 32): distance expansion per atom pair
101
+ ```
102
+ """
103
+
104
+ def __init__(self, num_kernel: int = 128, edge_types: int = 4096, **kwargs):
105
+ super().__init__(**kwargs)
106
+ self.num_kernel = num_kernel
107
+ self.edge_types = edge_types
108
+
109
+ self.mul = layers.Embedding(input_dim=edge_types, output_dim=1, name="mul")
110
+ self.bias = layers.Embedding(input_dim=edge_types, output_dim=1, name="bias")
111
+
112
+ def build(self, input_shape=None):
113
+ if not self.built:
114
+ self.means = self.add_weight(
115
+ shape=(self.num_kernel,),
116
+ initializer="uniform",
117
+ trainable=True,
118
+ name="means",
119
+ )
120
+ self.temps = self.add_weight(
121
+ shape=(self.num_kernel,),
122
+ initializer="uniform",
123
+ trainable=True,
124
+ name="temps",
125
+ )
126
+ super().build(input_shape)
127
+
128
+ def call(self, dist, edge_types):
129
+ mul = self.mul(edge_types)
130
+ bias = self.bias(edge_types)
131
+ x = mul * ops.expand_dims(dist, axis=-1) + bias
132
+ means = self.means
133
+ temps = ops.abs(self.temps)
134
+ return ops.exp(-temps * ops.power(x - means, 2))
135
+
136
+
137
+ class NonLinear(layers.Layer):
138
+ r"""Two-layer MLP with GELU activation.
139
+
140
+ Args:
141
+ hidden_dim (int): Intermediate hidden dimension.
142
+ output_dim (int): Output dimension.
143
+ **kwargs: Additional layer arguments.
144
+ """
145
+
146
+ def __init__(self, hidden_dim: int, output_dim: int, **kwargs):
147
+ super().__init__(**kwargs)
148
+ self.hidden_dim = hidden_dim
149
+ self.output_dim = output_dim
150
+
151
+ self.layer1 = layers.Dense(hidden_dim, name="layer1")
152
+ self.layer2 = layers.Dense(output_dim, name="layer2")
153
+
154
+ def build(self, input_shape=None):
155
+ if not self.built:
156
+ self.layer1.build((None, self.hidden_dim))
157
+ self.layer2.build((None, self.hidden_dim))
158
+ super().build(input_shape)
159
+
160
+ def call(self, x):
161
+ return self.layer2(ops.gelu(self.layer1(x)))
162
+
163
+
164
+ class SelfMultiheadAttention(layers.Layer):
165
+ r"""Fused query-key-value self-attention with additive attention bias for 3D Graphormer.
166
+
167
+ Args:
168
+ embed_dim (int): Embedding dimension.
169
+ num_heads (int): Number of attention heads.
170
+ dropout (float, optional): Attention dropout. (default: ``0.0``)
171
+ bias (bool, optional): Whether to use projection biases. (default: ``True``)
172
+ scaling_factor (float, optional): Attention scaling factor. (default: ``1.0``)
173
+ **kwargs: Additional layer arguments.
174
+ """
175
+
176
+ def __init__(
177
+ self,
178
+ embed_dim: int,
179
+ num_heads: int,
180
+ dropout: float = 0.0,
181
+ bias: bool = True,
182
+ scaling_factor: float = 1.0,
183
+ **kwargs,
184
+ ):
185
+ super().__init__(**kwargs)
186
+ self.embed_dim = embed_dim
187
+ self.num_heads = num_heads
188
+ self.head_dim = embed_dim // num_heads
189
+ self.dropout_rate = dropout
190
+ self.scaling = (self.head_dim * scaling_factor) ** -0.5
191
+
192
+ self.in_proj = layers.Dense(embed_dim * 3, use_bias=bias, name="in_proj")
193
+ self.out_proj = layers.Dense(embed_dim, use_bias=bias, name="out_proj")
194
+ self.dropout = layers.Dropout(dropout)
195
+
196
+ def build(self, input_shape=None):
197
+ if not self.built:
198
+ self.in_proj.build((None, None, self.embed_dim))
199
+ self.out_proj.build((None, None, self.embed_dim))
200
+ super().build(input_shape)
201
+
202
+ def call(self, x, attn_bias=None, training: bool = False):
203
+ r"""
204
+ Args:
205
+ x (Tensor): Node sequence tensor of shape ``[batch_size, num_nodes, embed_dim]``.
206
+ attn_bias (Tensor, optional): Bias of shape ``[batch_size * num_heads, num_nodes, num_nodes]``.
207
+ training (bool, optional): Training flag. (default: ``False``)
208
+
209
+ Returns:
210
+ Tensor: Attention output of shape ``[batch_size, num_nodes, embed_dim]``.
211
+ """
212
+ shape = ops.shape(x)
213
+ batch_size, num_nodes = shape[0], shape[1]
214
+
215
+ qkv = self.in_proj(x)
216
+ # Split into q, k, v each of shape [batch_size, num_nodes, embed_dim]
217
+ q = qkv[:, :, : self.embed_dim]
218
+ k = qkv[:, :, self.embed_dim : 2 * self.embed_dim]
219
+ v = qkv[:, :, 2 * self.embed_dim :]
220
+
221
+ # Reshape to [batch_size * num_heads, num_nodes, head_dim]
222
+ q = ops.reshape(
223
+ ops.transpose(
224
+ ops.reshape(q, (batch_size, num_nodes, self.num_heads, self.head_dim)),
225
+ (0, 2, 1, 3),
226
+ ),
227
+ (batch_size * self.num_heads, num_nodes, self.head_dim),
228
+ ) * self.scaling
229
+ k = ops.reshape(
230
+ ops.transpose(
231
+ ops.reshape(k, (batch_size, num_nodes, self.num_heads, self.head_dim)),
232
+ (0, 2, 1, 3),
233
+ ),
234
+ (batch_size * self.num_heads, num_nodes, self.head_dim),
235
+ )
236
+ v = ops.reshape(
237
+ ops.transpose(
238
+ ops.reshape(v, (batch_size, num_nodes, self.num_heads, self.head_dim)),
239
+ (0, 2, 1, 3),
240
+ ),
241
+ (batch_size * self.num_heads, num_nodes, self.head_dim),
242
+ )
243
+
244
+ attn_weights = ops.matmul(q, ops.transpose(k, (0, 2, 1)))
245
+ if attn_bias is not None:
246
+ attn_weights = attn_weights + attn_bias
247
+
248
+ attn_probs = ops.softmax(attn_weights, axis=-1)
249
+ attn_probs = self.dropout(attn_probs, training=training)
250
+
251
+ attn = ops.matmul(attn_probs, v)
252
+ # Reshape back to [batch_size, num_nodes, embed_dim]
253
+ attn = ops.reshape(
254
+ ops.transpose(
255
+ ops.reshape(attn, (batch_size, self.num_heads, num_nodes, self.head_dim)),
256
+ (0, 2, 1, 3),
257
+ ),
258
+ (batch_size, num_nodes, self.embed_dim),
259
+ )
260
+ return self.out_proj(attn)
261
+
262
+
263
+ class Graphormer3DEncoderLayer(layers.Layer):
264
+ r"""3D Graphormer Transformer Encoder Layer with Pre-LN.
265
+
266
+ Args:
267
+ embedding_dim (int): Embedding dimension.
268
+ ffn_embedding_dim (int): FFN hidden dimension.
269
+ num_attention_heads (int): Number of attention heads.
270
+ dropout (float, optional): Dropout probability. (default: ``0.1``)
271
+ attention_dropout (float, optional): Attention dropout. (default: ``0.1``)
272
+ activation_dropout (float, optional): Activation dropout. (default: ``0.1``)
273
+ **kwargs: Additional layer arguments.
274
+
275
+ Example:
276
+ ```python
277
+ import numpy as np
278
+ from k3_node.models import Graphormer3DEncoderLayer
279
+
280
+ x = np.random.rand(2, 4, 32).astype("float32")
281
+ attn_bias = np.zeros((2 * 4, 4, 4), dtype="float32") # [batch * heads, atoms, atoms]
282
+ layer = Graphormer3DEncoderLayer(embedding_dim=32, ffn_embedding_dim=64, num_attention_heads=4)
283
+ print(tuple(layer(x, attn_bias=attn_bias).shape)) # (2, 4, 32)
284
+ ```
285
+ """
286
+
287
+ def __init__(
288
+ self,
289
+ embedding_dim: int = 768,
290
+ ffn_embedding_dim: int = 3072,
291
+ num_attention_heads: int = 8,
292
+ dropout: float = 0.1,
293
+ attention_dropout: float = 0.1,
294
+ activation_dropout: float = 0.1,
295
+ **kwargs,
296
+ ):
297
+ super().__init__(**kwargs)
298
+ self.embedding_dim = embedding_dim
299
+ self.ffn_embedding_dim = ffn_embedding_dim
300
+ self.num_attention_heads = num_attention_heads
301
+ self.dropout_rate = dropout
302
+ self.activation_dropout_rate = activation_dropout
303
+
304
+ self.self_attn = SelfMultiheadAttention(
305
+ embed_dim=embedding_dim,
306
+ num_heads=num_attention_heads,
307
+ dropout=attention_dropout,
308
+ name="self_attn",
309
+ )
310
+ self.self_attn_layer_norm = layers.LayerNormalization(
311
+ epsilon=1e-5, name="self_attn_layer_norm"
312
+ )
313
+ self.fc1 = layers.Dense(ffn_embedding_dim, name="fc1")
314
+ self.fc2 = layers.Dense(embedding_dim, name="fc2")
315
+ self.final_layer_norm = layers.LayerNormalization(
316
+ epsilon=1e-5, name="final_layer_norm"
317
+ )
318
+
319
+ self.dropout = layers.Dropout(dropout)
320
+ self.act_dropout = layers.Dropout(activation_dropout)
321
+
322
+ def build(self, input_shape=None):
323
+ if not self.built:
324
+ self.self_attn.build((None, None, self.embedding_dim))
325
+ self.self_attn_layer_norm.build((None, None, self.embedding_dim))
326
+ self.fc1.build((None, None, self.embedding_dim))
327
+ self.fc2.build((None, None, self.ffn_embedding_dim))
328
+ self.final_layer_norm.build((None, None, self.embedding_dim))
329
+ super().build(input_shape)
330
+
331
+ def call(self, x, attn_bias=None, training: bool = False):
332
+ residual = x
333
+ x = self.self_attn_layer_norm(x)
334
+ x = self.self_attn(x, attn_bias=attn_bias, training=training)
335
+ x = self.dropout(x, training=training)
336
+ x = residual + x
337
+
338
+ residual = x
339
+ x = self.final_layer_norm(x)
340
+ x = ops.gelu(self.fc1(x))
341
+ x = self.act_dropout(x, training=training)
342
+ x = self.fc2(x)
343
+ x = self.dropout(x, training=training)
344
+ x = residual + x
345
+ return x
346
+
347
+
348
+ class NodeTaskHead(layers.Layer):
349
+ r"""Rotational-equivariant 3D vector force prediction head.
350
+
351
+ Args:
352
+ embed_dim (int): Embedding dimension.
353
+ num_heads (int): Number of attention heads.
354
+ **kwargs: Additional layer arguments.
355
+
356
+ Example:
357
+ ```python
358
+ import numpy as np
359
+ from k3_node.models import NodeTaskHead
360
+
361
+ query = np.random.rand(2, 5, 32).astype("float32") # atom representations
362
+ attn_bias = np.zeros((2 * 4, 5, 5), dtype="float32")
363
+ delta_pos = np.random.rand(2, 5, 5, 3).astype("float32") # pairwise displacement vectors
364
+ head = NodeTaskHead(embed_dim=32, num_heads=4)
365
+ print(tuple(head(query, attn_bias, delta_pos).shape)) # (2, 5, 3): a force vector per atom
366
+ ```
367
+ """
368
+
369
+ def __init__(self, embed_dim: int, num_heads: int, **kwargs):
370
+ super().__init__(**kwargs)
371
+ self.embed_dim = embed_dim
372
+ self.num_heads = num_heads
373
+ self.head_dim = embed_dim // num_heads
374
+ self.scaling = self.head_dim ** -0.5
375
+
376
+ self.q_proj = layers.Dense(embed_dim, name="q_proj")
377
+ self.k_proj = layers.Dense(embed_dim, name="k_proj")
378
+ self.v_proj = layers.Dense(embed_dim, name="v_proj")
379
+
380
+ self.force_proj1 = layers.Dense(1, name="force_proj1")
381
+ self.force_proj2 = layers.Dense(1, name="force_proj2")
382
+ self.force_proj3 = layers.Dense(1, name="force_proj3")
383
+
384
+ def build(self, input_shape=None):
385
+ if not self.built:
386
+ self.q_proj.build((None, None, self.embed_dim))
387
+ self.k_proj.build((None, None, self.embed_dim))
388
+ self.v_proj.build((None, None, self.embed_dim))
389
+ self.force_proj1.build((None, None, self.embed_dim))
390
+ self.force_proj2.build((None, None, self.embed_dim))
391
+ self.force_proj3.build((None, None, self.embed_dim))
392
+ super().build(input_shape)
393
+
394
+ def call(self, query, attn_bias, delta_pos, training: bool = False):
395
+ r"""
396
+ Args:
397
+ query (Tensor): Node representation of shape ``[batch_size, num_nodes, embed_dim]``.
398
+ attn_bias (Tensor): Attention bias of shape ``[batch_size * num_heads, num_nodes, num_nodes]``.
399
+ delta_pos (Tensor): Normalized unit direction vectors of shape ``[batch_size, num_nodes, num_nodes, 3]``.
400
+ training (bool, optional): Training flag. (default: ``False``)
401
+
402
+ Returns:
403
+ Tensor: Predicted 3D forces of shape ``[batch_size, num_nodes, 3]``.
404
+ """
405
+ shape = ops.shape(query)
406
+ bsz, n_node = shape[0], shape[1]
407
+
408
+ q = self.q_proj(query) * self.scaling
409
+ k = self.k_proj(query)
410
+ v = self.v_proj(query)
411
+
412
+ # Reshape to [bsz, num_heads, n_node, head_dim]
413
+ q = ops.transpose(ops.reshape(q, (bsz, n_node, self.num_heads, self.head_dim)), (0, 2, 1, 3))
414
+ k = ops.transpose(ops.reshape(k, (bsz, n_node, self.num_heads, self.head_dim)), (0, 2, 1, 3))
415
+ v = ops.transpose(ops.reshape(v, (bsz, n_node, self.num_heads, self.head_dim)), (0, 2, 1, 3))
416
+
417
+ attn = ops.matmul(q, ops.transpose(k, (0, 1, 3, 2))) # [bsz, num_heads, n_node, n_node]
418
+ attn_flat = ops.reshape(attn, (-1, n_node, n_node)) + attn_bias
419
+ attn_probs = ops.softmax(attn_flat, axis=-1)
420
+ attn_probs = ops.reshape(attn_probs, (bsz, self.num_heads, n_node, n_node))
421
+
422
+ # rot_attn_probs: [bsz, num_heads, n_node, n_node, 3]
423
+ rot_attn_probs = ops.expand_dims(attn_probs, axis=-1) * ops.expand_dims(delta_pos, axis=1)
424
+ # Permute to [bsz, num_heads, 3, n_node, n_node]
425
+ rot_attn_probs = ops.transpose(rot_attn_probs, (0, 1, 4, 2, 3))
426
+
427
+ # Multiply with v: [bsz, num_heads, 1, n_node, head_dim]
428
+ v_exp = ops.expand_dims(v, axis=2)
429
+ x = ops.matmul(rot_attn_probs, v_exp) # [bsz, num_heads, 3, n_node, head_dim]
430
+
431
+ # Permute to [bsz, n_node, 3, num_heads, head_dim] -> [bsz, n_node, 3, embed_dim]
432
+ x = ops.transpose(x, (0, 3, 2, 1, 4))
433
+ x = ops.reshape(x, (bsz, n_node, 3, self.embed_dim))
434
+
435
+ f1 = self.force_proj1(x[:, :, 0, :]) # [bsz, n_node, 1]
436
+ f2 = self.force_proj2(x[:, :, 1, :]) # [bsz, n_node, 1]
437
+ f3 = self.force_proj3(x[:, :, 2, :]) # [bsz, n_node, 1]
438
+
439
+ cur_force = ops.concatenate([f1, f2, f3], axis=-1) # [bsz, n_node, 3]
440
+ return cur_force
441
+
442
+
443
+ class Graphormer3D(keras.Model):
444
+ r"""Graphormer-3D model for 3D molecular structure modeling, energy, and force prediction
445
+ from `"Benchmarking Graphormer on Large-Scale Molecular Modeling Datasets" <https://arxiv.org/abs/2203.04810>`_.
446
+
447
+ Args:
448
+ layers (int, optional): Number of encoder layers per block. (default: ``12``)
449
+ blocks (int, optional): Number of repeated encoder blocks. (default: ``4``)
450
+ embed_dim (int, optional): Hidden embedding dimension. (default: ``768``)
451
+ ffn_embed_dim (int, optional): FFN intermediate dimension. (default: ``768``)
452
+ attention_heads (int, optional): Number of attention heads. (default: ``48``)
453
+ num_kernel (int, optional): Number of Gaussian basis kernels. (default: ``128``)
454
+ atom_types (int, optional): Number of atom types. (default: ``64``)
455
+ dropout (float, optional): Dropout probability. (default: ``0.1``)
456
+ attention_dropout (float, optional): Attention dropout. (default: ``0.1``)
457
+ activation_dropout (float, optional): FFN activation dropout. (default: ``0.0``)
458
+ input_dropout (float, optional): Input features dropout. (default: ``0.0``)
459
+ **kwargs: Additional model arguments.
460
+
461
+ Example:
462
+ ```python
463
+ import numpy as np
464
+ from k3_node.models import Graphormer3D
465
+
466
+ atoms = np.array([[1, 2, 3, 4, 0], [2, 3, 4, 0, 0]]) # atom types, 0 = padding
467
+ tags = np.array([[1, 1, 2, 2, 0], [1, 2, 2, 0, 0]]) # e.g. surface/adsorbate tags (OC20)
468
+ pos = np.random.rand(2, 5, 3).astype("float32")
469
+ model = Graphormer3D(layers=2, blocks=2, embed_dim=32, ffn_embed_dim=64, attention_heads=4,
470
+ num_kernel=16, atom_types=16)
471
+ energy, forces = model(atoms, tags, pos)
472
+ print(tuple(energy.shape), tuple(forces.shape)) # (2,) (2, 5, 3)
473
+ ```
474
+ """
475
+
476
+ def __init__(
477
+ self,
478
+ layers: int = 12,
479
+ blocks: int = 4,
480
+ embed_dim: int = 768,
481
+ ffn_embed_dim: int = 768,
482
+ attention_heads: int = 48,
483
+ num_kernel: int = 128,
484
+ atom_types: int = 64,
485
+ dropout: float = 0.1,
486
+ attention_dropout: float = 0.1,
487
+ activation_dropout: float = 0.0,
488
+ input_dropout: float = 0.0,
489
+ **kwargs,
490
+ ):
491
+ super().__init__(**kwargs)
492
+ self.num_encoder_layers = layers
493
+ self.blocks = blocks
494
+ self.embed_dim = embed_dim
495
+ self.ffn_embed_dim = ffn_embed_dim
496
+ self.attention_heads = attention_heads
497
+ self.num_kernel = num_kernel
498
+ self.atom_types = atom_types
499
+ self.edge_types = atom_types * atom_types
500
+
501
+ self.atom_encoder = keras.layers.Embedding(
502
+ input_dim=atom_types,
503
+ output_dim=embed_dim,
504
+ name="atom_encoder",
505
+ )
506
+ self.tag_encoder = keras.layers.Embedding(
507
+ input_dim=3,
508
+ output_dim=embed_dim,
509
+ name="tag_encoder",
510
+ )
511
+ self.input_dropout = keras.layers.Dropout(input_dropout)
512
+
513
+ self.encoder_layers = [
514
+ Graphormer3DEncoderLayer(
515
+ embedding_dim=embed_dim,
516
+ ffn_embedding_dim=ffn_embed_dim,
517
+ num_attention_heads=attention_heads,
518
+ dropout=dropout,
519
+ attention_dropout=attention_dropout,
520
+ activation_dropout=activation_dropout,
521
+ name=f"layers_{i}",
522
+ )
523
+ for i in range(layers)
524
+ ]
525
+
526
+ self.final_ln = keras.layers.LayerNormalization(epsilon=1e-5, name="final_ln")
527
+
528
+ self.energy_proj = NonLinear(embed_dim, 1, name="energy_proj")
529
+ self.energe_agg_factor = keras.layers.Embedding(
530
+ input_dim=3, output_dim=1, name="energe_agg_factor"
531
+ )
532
+
533
+ self.gbf = GaussianLayer(num_kernel, self.edge_types, name="gbf")
534
+ self.bias_proj = NonLinear(num_kernel, attention_heads, name="bias_proj")
535
+ self.edge_proj = keras.layers.Dense(embed_dim, name="edge_proj")
536
+ self.node_proc = NodeTaskHead(embed_dim, attention_heads, name="node_proc")
537
+
538
+ def build(self, input_shape=None):
539
+ if not self.built:
540
+ self.atom_encoder.build(None)
541
+ self.tag_encoder.build(None)
542
+ self.gbf.build(None)
543
+ self.bias_proj.build(None)
544
+ self.edge_proj.build((None, self.num_kernel))
545
+ for layer in self.encoder_layers:
546
+ layer.build((None, None, self.embed_dim))
547
+ self.final_ln.build((None, None, self.embed_dim))
548
+ self.energy_proj.build(None)
549
+ self.energe_agg_factor.build(None)
550
+ self.node_proc.build(None)
551
+ super().build(input_shape)
552
+
553
+ def call(
554
+ self,
555
+ atoms,
556
+ tags,
557
+ pos,
558
+ real_mask=None,
559
+ training: bool = False,
560
+ ):
561
+ r"""Forward pass for Graphormer-3D predicting total energy and atomic forces.
562
+
563
+ Args:
564
+ atoms (Tensor): Atom indices of shape ``[batch_size, num_nodes]``.
565
+ tags (Tensor): Tag indices of shape ``[batch_size, num_nodes]`` (0: fixed, 1: sub-surface, 2: surface).
566
+ pos (Tensor): 3D atomic coordinates of shape ``[batch_size, num_nodes, 3]``.
567
+ real_mask (Tensor, optional): Valid non-padding mask of shape ``[batch_size, num_nodes]``.
568
+ If None, non-zero atom indices are considered valid.
569
+ training (bool, optional): Training flag. (default: ``False``)
570
+
571
+ Returns:
572
+ Tuple[Tensor, Tensor]: Tuple of predicted energy ``[batch_size]`` and atomic forces
573
+ ``[batch_size, num_nodes, 3]``.
574
+ """
575
+ shape = ops.shape(atoms)
576
+ n_graph, n_node = shape[0], shape[1]
577
+
578
+ padding_mask = ops.equal(atoms, 0)
579
+ if real_mask is None:
580
+ real_mask = ops.logical_not(padding_mask)
581
+
582
+ # Pairwise displacement vectors and Euclidean distances
583
+ delta_pos = ops.expand_dims(pos, axis=1) - ops.expand_dims(pos, axis=2) # [B, N, N, 3]
584
+ dist = ops.sqrt(ops.sum(ops.power(delta_pos, 2), axis=-1) + 1e-12) # [B, N, N]
585
+ norm_delta_pos = delta_pos / (ops.expand_dims(dist, axis=-1) + 1e-5)
586
+
587
+ # Edge type indices: [B, N, N]
588
+ edge_type = (
589
+ ops.expand_dims(atoms, axis=2) * self.atom_types
590
+ + ops.expand_dims(atoms, axis=1)
591
+ )
592
+
593
+ gbf_feature = self.gbf(dist, edge_type) # [B, N, N, K]
594
+
595
+ # Mask padding in edge features
596
+ pad_edge_mask = ops.expand_dims(ops.expand_dims(padding_mask, axis=1), axis=-1)
597
+ edge_features = ops.where(pad_edge_mask, 0.0, gbf_feature)
598
+
599
+ graph_node_feature = (
600
+ self.tag_encoder(tags)
601
+ + self.atom_encoder(atoms)
602
+ + self.edge_proj(ops.sum(edge_features, axis=-2))
603
+ )
604
+
605
+ output = self.input_dropout(graph_node_feature, training=training)
606
+
607
+ # Attention bias: [B, N, N, num_heads] -> [B, num_heads, N, N]
608
+ graph_attn_bias = ops.transpose(self.bias_proj(gbf_feature), (0, 3, 1, 2))
609
+ # Mask padding: [B, 1, 1, N]
610
+ pad_mask_bias = ops.expand_dims(ops.expand_dims(padding_mask, axis=1), axis=2)
611
+ graph_attn_bias = ops.where(pad_mask_bias, -1e9, graph_attn_bias)
612
+ graph_attn_bias = ops.reshape(graph_attn_bias, (-1, n_node, n_node))
613
+
614
+ # Multi-block, multi-layer Transformer Encoder
615
+ for _ in range(self.blocks):
616
+ for enc_layer in self.encoder_layers:
617
+ output = enc_layer(output, attn_bias=graph_attn_bias, training=training)
618
+
619
+ output = self.final_ln(output)
620
+
621
+ # Energy prediction
622
+ eng_output = ops.squeeze(
623
+ self.energy_proj(output) * self.energe_agg_factor(tags), axis=-1
624
+ )
625
+ output_mask = ops.logical_and(tags > 0, real_mask)
626
+ eng_output = ops.sum(ops.where(output_mask, eng_output, 0.0), axis=-1)
627
+
628
+ # Force prediction
629
+ node_output = self.node_proc(
630
+ output, graph_attn_bias, norm_delta_pos, training=training
631
+ )
632
+
633
+ return eng_output, node_output
634
+
635
+ @classmethod
636
+ def from_pretrained(
637
+ cls,
638
+ pretrained_name: str = "oc20is2re_graphormer3d_base",
639
+ folder: str = "checkpoints",
640
+ download: bool = True,
641
+ **kwargs,
642
+ ) -> "Graphormer3D":
643
+ r"""Instantiates a Graphormer3D model with pre-trained weights."""
644
+ cfg = get_graphormer3d_config(pretrained_name)
645
+ cfg.update(kwargs)
646
+ model = cls(**cfg)
647
+ load_graphormer3d_weights(model, pretrained_name=pretrained_name, folder=folder, download=download)
648
+ return model
649
+
650
+ def __repr__(self) -> str:
651
+ return (
652
+ f"{self.__class__.__name__}("
653
+ f"blocks={self.blocks}, "
654
+ f"layers={self.num_encoder_layers}, "
655
+ f"embed_dim={self.embed_dim}, "
656
+ f"attention_heads={self.attention_heads}, "
657
+ f"num_kernel={self.num_kernel})"
658
+ )
659
+
660
+
661
+ PRETRAINED_3D_URLS = {
662
+ "oc20is2re_graphormer3d_base": "https://szheng.blob.core.windows.net/graphormer/modelzoo/oc20is2re/checkpoint_last_oc20_is2re.pt",
663
+ }
664
+
665
+
666
+ def get_graphormer3d_config(name_or_variant: str) -> Dict[str, Any]:
667
+ r"""Returns configuration dictionary for Graphormer-3D."""
668
+ return {
669
+ "blocks": 4,
670
+ "layers": 12,
671
+ "embed_dim": 768,
672
+ "ffn_embed_dim": 768,
673
+ "attention_heads": 48,
674
+ "num_kernel": 128,
675
+ "atom_types": 64,
676
+ "dropout": 0.1,
677
+ "attention_dropout": 0.1,
678
+ "activation_dropout": 0.0,
679
+ }
680
+
681
+
682
+ def download_graphormer3d_checkpoint(
683
+ name: str = "oc20is2re_graphormer3d_base",
684
+ folder: str = "checkpoints",
685
+ log: bool = True,
686
+ ) -> str:
687
+ r"""Downloads a pre-trained Graphormer-3D checkpoint."""
688
+ from k3_node.data.download import download_url
689
+
690
+ clean_name = name.lower().replace("-", "_")
691
+ if clean_name not in PRETRAINED_3D_URLS and name not in PRETRAINED_3D_URLS:
692
+ raise ValueError(
693
+ f"Unknown pretrained 3D model name '{name}'. Available: {list(PRETRAINED_3D_URLS.keys())}"
694
+ )
695
+
696
+ matched_key = clean_name if clean_name in PRETRAINED_3D_URLS else name
697
+ url = PRETRAINED_3D_URLS[matched_key]
698
+ filename = f"{matched_key}.pt"
699
+ local_path = os.path.join(folder, filename)
700
+
701
+ if os.path.exists(local_path):
702
+ return local_path
703
+
704
+ repo_alt = os.path.join("Graphormer-main", "checkpoints", filename)
705
+ if os.path.exists(repo_alt):
706
+ return repo_alt
707
+
708
+ return download_url(url, folder=folder, filename=filename, log=log)
709
+
710
+
711
+ def load_graphormer3d_weights(
712
+ model: Graphormer3D,
713
+ checkpoint_path: Optional[str] = None,
714
+ pretrained_name: Optional[str] = None,
715
+ folder: str = "checkpoints",
716
+ download: bool = True,
717
+ ) -> Graphormer3D:
718
+ r"""Loads weights from a PyTorch checkpoint into a Graphormer3D model."""
719
+ path_to_load = checkpoint_path
720
+
721
+ if path_to_load is None:
722
+ if pretrained_name is None:
723
+ raise ValueError("Either checkpoint_path or pretrained_name must be specified.")
724
+ candidate = os.path.join(folder, f"{pretrained_name}.pt")
725
+ if os.path.isfile(candidate):
726
+ path_to_load = candidate
727
+ elif download:
728
+ path_to_load = download_graphormer3d_checkpoint(pretrained_name, folder=folder)
729
+ else:
730
+ raise FileNotFoundError(f"Checkpoint for '{pretrained_name}' not found at '{candidate}'.")
731
+ elif not os.path.isfile(path_to_load) and download and pretrained_name:
732
+ path_to_load = download_graphormer3d_checkpoint(pretrained_name, folder=folder)
733
+
734
+ import torch
735
+ import numpy as np
736
+
737
+ state = torch.load(path_to_load, map_location="cpu")
738
+ if isinstance(state, dict) and "model" in state:
739
+ state_dict = state["model"]
740
+ elif isinstance(state, dict):
741
+ state_dict = state
742
+ else:
743
+ raise ValueError(f"Unexpected checkpoint format: {type(state)}")
744
+
745
+ if not model.built:
746
+ model.build(None)
747
+
748
+ def _to_tensor(t):
749
+ if hasattr(t, "detach"):
750
+ t = t.detach()
751
+ if hasattr(t, "numpy"):
752
+ t = t.numpy()
753
+ return ops.convert_to_tensor(np.array(t, dtype=np.float32), dtype="float32")
754
+
755
+ clean_dict = {}
756
+ for k, v in state_dict.items():
757
+ ck = k
758
+ if ck.startswith("encoder."):
759
+ ck = ck[len("encoder.") :]
760
+ clean_dict[ck] = v
761
+
762
+ # Atom & Tag embeddings
763
+ if "atom_encoder.weight" in clean_dict:
764
+ model.atom_encoder.weights[0].assign(_to_tensor(clean_dict["atom_encoder.weight"]))
765
+ if "tag_encoder.weight" in clean_dict:
766
+ model.tag_encoder.weights[0].assign(_to_tensor(clean_dict["tag_encoder.weight"]))
767
+
768
+ # GBF
769
+ for name in ["means", "stds", "mul", "bias"]:
770
+ key = f"gbf.{name}.weight"
771
+ if key in clean_dict:
772
+ getattr(model.gbf, name).weights[0].assign(_to_tensor(clean_dict[key]))
773
+
774
+ # Bias projection
775
+ for layer_name in ["layer1", "layer2"]:
776
+ if f"bias_proj.{layer_name}.weight" in clean_dict:
777
+ getattr(model.bias_proj, layer_name).kernel.assign(
778
+ _to_tensor(clean_dict[f"bias_proj.{layer_name}.weight"].t())
779
+ )
780
+ if f"bias_proj.{layer_name}.bias" in clean_dict:
781
+ getattr(model.bias_proj, layer_name).bias.assign(
782
+ _to_tensor(clean_dict[f"bias_proj.{layer_name}.bias"])
783
+ )
784
+
785
+ # Edge projection
786
+ if "edge_proj.weight" in clean_dict:
787
+ model.edge_proj.kernel.assign(_to_tensor(clean_dict["edge_proj.weight"].t()))
788
+ if "edge_proj.bias" in clean_dict:
789
+ model.edge_proj.bias.assign(_to_tensor(clean_dict["edge_proj.bias"]))
790
+
791
+ # Encoder layers
792
+ for i, enc_layer in enumerate(model.encoder_layers):
793
+ p = f"layers.{i}"
794
+ # Self attention in_proj & out_proj
795
+ if f"{p}.self_attn.in_proj.weight" in clean_dict:
796
+ enc_layer.self_attn.in_proj.kernel.assign(
797
+ _to_tensor(clean_dict[f"{p}.self_attn.in_proj.weight"].t())
798
+ )
799
+ if f"{p}.self_attn.in_proj.bias" in clean_dict:
800
+ enc_layer.self_attn.in_proj.bias.assign(
801
+ _to_tensor(clean_dict[f"{p}.self_attn.in_proj.bias"])
802
+ )
803
+ if f"{p}.self_attn.out_proj.weight" in clean_dict:
804
+ enc_layer.self_attn.out_proj.kernel.assign(
805
+ _to_tensor(clean_dict[f"{p}.self_attn.out_proj.weight"].t())
806
+ )
807
+ if f"{p}.self_attn.out_proj.bias" in clean_dict:
808
+ enc_layer.self_attn.out_proj.bias.assign(
809
+ _to_tensor(clean_dict[f"{p}.self_attn.out_proj.bias"])
810
+ )
811
+
812
+ # Norms
813
+ if f"{p}.self_attn_layer_norm.weight" in clean_dict:
814
+ enc_layer.self_attn_layer_norm.gamma.assign(
815
+ _to_tensor(clean_dict[f"{p}.self_attn_layer_norm.weight"])
816
+ )
817
+ if f"{p}.self_attn_layer_norm.bias" in clean_dict:
818
+ enc_layer.self_attn_layer_norm.beta.assign(
819
+ _to_tensor(clean_dict[f"{p}.self_attn_layer_norm.bias"])
820
+ )
821
+ if f"{p}.final_layer_norm.weight" in clean_dict:
822
+ enc_layer.final_layer_norm.gamma.assign(
823
+ _to_tensor(clean_dict[f"{p}.final_layer_norm.weight"])
824
+ )
825
+ if f"{p}.final_layer_norm.bias" in clean_dict:
826
+ enc_layer.final_layer_norm.beta.assign(
827
+ _to_tensor(clean_dict[f"{p}.final_layer_norm.bias"])
828
+ )
829
+
830
+ # FFN
831
+ if f"{p}.fc1.weight" in clean_dict:
832
+ enc_layer.fc1.kernel.assign(_to_tensor(clean_dict[f"{p}.fc1.weight"].t()))
833
+ if f"{p}.fc1.bias" in clean_dict:
834
+ enc_layer.fc1.bias.assign(_to_tensor(clean_dict[f"{p}.fc1.bias"]))
835
+ if f"{p}.fc2.weight" in clean_dict:
836
+ enc_layer.fc2.kernel.assign(_to_tensor(clean_dict[f"{p}.fc2.weight"].t()))
837
+ if f"{p}.fc2.bias" in clean_dict:
838
+ enc_layer.fc2.bias.assign(_to_tensor(clean_dict[f"{p}.fc2.bias"]))
839
+
840
+ # Final LayerNorm
841
+ if "final_ln.weight" in clean_dict:
842
+ model.final_ln.gamma.assign(_to_tensor(clean_dict["final_ln.weight"]))
843
+ if "final_ln.bias" in clean_dict:
844
+ model.final_ln.beta.assign(_to_tensor(clean_dict["final_ln.bias"]))
845
+
846
+ # Energy head
847
+ for layer_name in ["layer1", "layer2"]:
848
+ if f"engergy_proj.{layer_name}.weight" in clean_dict:
849
+ getattr(model.energy_proj, layer_name).kernel.assign(
850
+ _to_tensor(clean_dict[f"engergy_proj.{layer_name}.weight"].t())
851
+ )
852
+ if f"engergy_proj.{layer_name}.bias" in clean_dict:
853
+ getattr(model.energy_proj, layer_name).bias.assign(
854
+ _to_tensor(clean_dict[f"engergy_proj.{layer_name}.bias"])
855
+ )
856
+
857
+ if "energe_agg_factor.weight" in clean_dict:
858
+ model.energe_agg_factor.weights[0].assign(_to_tensor(clean_dict["energe_agg_factor.weight"]))
859
+
860
+ # NodeTaskHead (Forces)
861
+ for p_name in ["q_proj", "k_proj", "v_proj", "force_proj1", "force_proj2", "force_proj3"]:
862
+ proj = getattr(model.node_proc, p_name)
863
+ if f"node_proc.{p_name}.weight" in clean_dict:
864
+ proj.kernel.assign(_to_tensor(clean_dict[f"node_proc.{p_name}.weight"].t()))
865
+ if f"node_proc.{p_name}.bias" in clean_dict:
866
+ proj.bias.assign(_to_tensor(clean_dict[f"node_proc.{p_name}.bias"]))
867
+
868
+ return model