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,1258 @@
1
+ import math
2
+ import os
3
+ import urllib.request
4
+ from typing import Optional, Union, Tuple, List, Dict, Any
5
+
6
+ import keras
7
+ from keras import layers, ops
8
+ from k3_node.ops.creation import repeat
9
+
10
+
11
+ class GraphNodeFeature(layers.Layer):
12
+ r"""Computes initial node representations by summing atom feature embeddings,
13
+ in-degree embeddings, out-degree embeddings, and prepending a learnable graph token.
14
+
15
+ Args:
16
+ num_atoms (int): Maximum number of atom types.
17
+ num_in_degree (int): Maximum in-degree value.
18
+ num_out_degree (int): Maximum out-degree value.
19
+ hidden_dim (int): Embedding dimension.
20
+ **kwargs: Additional layer arguments.
21
+
22
+ Example:
23
+ ```python
24
+ import numpy as np
25
+ from k3_node.models import GraphNodeFeature
26
+
27
+ x = np.random.randint(1, 16, size=(2, 5)) # atom types of 2 graphs with 5 nodes
28
+ in_degree = np.random.randint(0, 10, size=(2, 5))
29
+ out_degree = np.random.randint(0, 10, size=(2, 5))
30
+ layer = GraphNodeFeature(num_atoms=16, num_in_degree=10, num_out_degree=10, hidden_dim=32)
31
+ print(tuple(layer(x, in_degree, out_degree).shape)) # (2, 6, 32): nodes + a virtual graph token
32
+ ```
33
+ """
34
+
35
+ def __init__(
36
+ self,
37
+ num_atoms: int,
38
+ num_in_degree: int,
39
+ num_out_degree: int,
40
+ hidden_dim: int,
41
+ **kwargs,
42
+ ):
43
+ super().__init__(**kwargs)
44
+ self.num_atoms = num_atoms
45
+ self.num_in_degree = num_in_degree
46
+ self.num_out_degree = num_out_degree
47
+ self.hidden_dim = hidden_dim
48
+
49
+ # 1-indexed padding_idx=0 in fairseq / Graphormer
50
+ self.atom_encoder = layers.Embedding(
51
+ input_dim=num_atoms + 1,
52
+ output_dim=hidden_dim,
53
+ mask_zero=False,
54
+ name="atom_encoder",
55
+ )
56
+ self.in_degree_encoder = layers.Embedding(
57
+ input_dim=num_in_degree,
58
+ output_dim=hidden_dim,
59
+ mask_zero=False,
60
+ name="in_degree_encoder",
61
+ )
62
+ self.out_degree_encoder = layers.Embedding(
63
+ input_dim=num_out_degree,
64
+ output_dim=hidden_dim,
65
+ mask_zero=False,
66
+ name="out_degree_encoder",
67
+ )
68
+ self.graph_token = layers.Embedding(
69
+ input_dim=1,
70
+ output_dim=hidden_dim,
71
+ name="graph_token",
72
+ )
73
+
74
+ def build(self, input_shape=None):
75
+ if not self.built:
76
+ self.atom_encoder.build(None)
77
+ self.in_degree_encoder.build(None)
78
+ self.out_degree_encoder.build(None)
79
+ self.graph_token.build(None)
80
+ super().build(input_shape)
81
+
82
+ def call(self, x, in_degree, out_degree):
83
+ r"""
84
+ Args:
85
+ x (Tensor): Atom features of shape ``[batch_size, num_nodes, num_features]``
86
+ or ``[batch_size, num_nodes]``.
87
+ in_degree (Tensor): In-degrees of shape ``[batch_size, num_nodes]``.
88
+ out_degree (Tensor): Out-degrees of shape ``[batch_size, num_nodes]``.
89
+
90
+ Returns:
91
+ Tensor: Node features with prepended graph token of shape
92
+ ``[batch_size, num_nodes + 1, hidden_dim]``.
93
+ """
94
+ x_shape = ops.shape(x)
95
+ batch_size, num_nodes = x_shape[0], x_shape[1]
96
+
97
+ if len(ops.shape(x)) == 2:
98
+ x_emb = self.atom_encoder(x)
99
+ else:
100
+ # Multi-dimensional atom features: sum over feature dim
101
+ x_emb = ops.sum(self.atom_encoder(x), axis=-2)
102
+
103
+ node_feature = (
104
+ x_emb
105
+ + self.in_degree_encoder(in_degree)
106
+ + self.out_degree_encoder(out_degree)
107
+ )
108
+
109
+ # Graph token feature: shape [batch_size, 1, hidden_dim]
110
+ token_id = ops.zeros((batch_size, 1), dtype="int32")
111
+ graph_token_feature = self.graph_token(token_id)
112
+
113
+ graph_node_feature = ops.concatenate([graph_token_feature, node_feature], axis=1)
114
+ return graph_node_feature
115
+
116
+
117
+ class GraphAttnBias(layers.Layer):
118
+ r"""Computes the structural attention bias for each attention head from shortest path
119
+ distances (spatial encoding) and edge features (edge encoding).
120
+
121
+ Args:
122
+ num_heads (int): Number of attention heads.
123
+ num_atoms (int): Maximum number of atom types.
124
+ num_edges (int): Maximum number of edge types.
125
+ num_spatial (int): Maximum spatial distance.
126
+ num_edge_dis (int): Maximum edge distance multiplier for multi-hop encoding.
127
+ edge_type (str, optional): Type of edge encoding (``"multi_hop"`` or ``"single_hop"``).
128
+ (default: ``"multi_hop"``)
129
+ multi_hop_max_dist (int, optional): Maximum distance for multi-hop paths. (default: ``20``)
130
+ **kwargs: Additional layer arguments.
131
+
132
+ Example:
133
+ ```python
134
+ import numpy as np
135
+ from k3_node.models import GraphAttnBias
136
+
137
+ attn_bias = np.zeros((2, 5, 5), dtype="float32") # 2 graphs, 4 nodes + graph token
138
+ spatial_pos = np.random.randint(0, 10, size=(2, 4, 4)) # shortest-path distances
139
+ x = np.zeros((2, 4, 1), dtype="int32")
140
+ edge_input = np.random.randint(0, 8, size=(2, 4, 4, 3, 2)) # edge types along each shortest path
141
+ layer = GraphAttnBias(num_heads=4, num_atoms=16, num_edges=8, num_spatial=10, num_edge_dis=5)
142
+ bias = layer(attn_bias=attn_bias, spatial_pos=spatial_pos, x=x, edge_input=edge_input)
143
+ print(tuple(bias.shape)) # (2, 4, 5, 5): one attention bias per head
144
+ ```
145
+ """
146
+
147
+ def __init__(
148
+ self,
149
+ num_heads: int,
150
+ num_atoms: int,
151
+ num_edges: int,
152
+ num_spatial: int,
153
+ num_edge_dis: int,
154
+ edge_type: str = "multi_hop",
155
+ multi_hop_max_dist: int = 20,
156
+ **kwargs,
157
+ ):
158
+ super().__init__(**kwargs)
159
+ self.num_heads = num_heads
160
+ self.num_atoms = num_atoms
161
+ self.num_edges = num_edges
162
+ self.num_spatial = num_spatial
163
+ self.num_edge_dis = num_edge_dis
164
+ self.edge_type = edge_type
165
+ self.multi_hop_max_dist = multi_hop_max_dist
166
+
167
+ self.edge_encoder = layers.Embedding(
168
+ input_dim=num_edges + 1,
169
+ output_dim=num_heads,
170
+ name="edge_encoder",
171
+ )
172
+ if self.edge_type == "multi_hop":
173
+ self.edge_dis_encoder = layers.Embedding(
174
+ input_dim=num_edge_dis * num_heads * num_heads,
175
+ output_dim=1,
176
+ name="edge_dis_encoder",
177
+ )
178
+ self.spatial_pos_encoder = layers.Embedding(
179
+ input_dim=num_spatial,
180
+ output_dim=num_heads,
181
+ name="spatial_pos_encoder",
182
+ )
183
+ self.graph_token_virtual_distance = layers.Embedding(
184
+ input_dim=1,
185
+ output_dim=num_heads,
186
+ name="graph_token_virtual_distance",
187
+ )
188
+
189
+ def build(self, input_shape=None):
190
+ if not self.built:
191
+ self.edge_encoder.build(None)
192
+ self.spatial_pos_encoder.build(None)
193
+ self.graph_token_virtual_distance.build(None)
194
+ if hasattr(self, "edge_dis_encoder"):
195
+ self.edge_dis_encoder.build(None)
196
+ super().build(input_shape)
197
+
198
+ def call(
199
+ self,
200
+ attn_bias,
201
+ spatial_pos,
202
+ x,
203
+ edge_input=None,
204
+ attn_edge_type=None,
205
+ ):
206
+ r"""
207
+ Args:
208
+ attn_bias (Tensor): Base attention mask of shape ``[batch_size, num_nodes + 1, num_nodes + 1]``.
209
+ spatial_pos (Tensor): Shortest path matrix of shape ``[batch_size, num_nodes, num_nodes]``.
210
+ x (Tensor): Atom features of shape ``[batch_size, num_nodes, ...]``.
211
+ edge_input (Tensor, optional): Multi-hop edge features along shortest paths of shape
212
+ ``[batch_size, num_nodes, num_nodes, max_dist, edge_feat_dim]``.
213
+ attn_edge_type (Tensor, optional): Single-hop edge types of shape
214
+ ``[batch_size, num_nodes, num_nodes, edge_feat_dim]``.
215
+
216
+ Returns:
217
+ Tensor: Attention bias tensor of shape ``[batch_size, num_heads, num_nodes + 1, num_nodes + 1]``.
218
+ """
219
+ x_shape = ops.shape(x)
220
+ batch_size, num_nodes = x_shape[0], x_shape[1]
221
+
222
+ # [batch_size, num_heads, num_nodes + 1, num_nodes + 1]
223
+ graph_attn_bias = repeat(
224
+ ops.expand_dims(attn_bias, axis=1), repeats=self.num_heads, axis=1
225
+ )
226
+
227
+ # Spatial position bias: [batch_size, num_nodes, num_nodes, num_heads] -> [batch_size, num_heads, num_nodes, num_nodes]
228
+ spatial_pos_bias = ops.transpose(self.spatial_pos_encoder(spatial_pos), (0, 3, 1, 2))
229
+
230
+ # Update node-to-node submatrix [1:, 1:]
231
+ sub_bias = graph_attn_bias[:, :, 1:, 1:] + spatial_pos_bias
232
+
233
+ # Virtual distance bias for graph token: shape [1, num_heads, 1]
234
+ t = ops.reshape(self.graph_token_virtual_distance(ops.zeros((1,), dtype="int32")), (1, self.num_heads, 1, 1))
235
+ # Add to row 0 (graph token to all nodes) and column 0 (all nodes to graph token)
236
+ row0 = graph_attn_bias[:, :, 0:1, 1:] + t[:, :, :, 0:]
237
+ col0 = graph_attn_bias[:, :, 1:, 0:1] + t[:, :, 0:, :]
238
+ corner = graph_attn_bias[:, :, 0:1, 0:1] + t
239
+
240
+ # Edge feature bias
241
+ if self.edge_type == "multi_hop" and edge_input is not None:
242
+ spatial_pos_ = ops.copy(spatial_pos)
243
+ # Replace 0 with 1 for padding
244
+ spatial_pos_ = ops.where(ops.equal(spatial_pos_, 0), 1, spatial_pos_)
245
+ spatial_pos_ = ops.where(spatial_pos_ > 1, spatial_pos_ - 1, spatial_pos_)
246
+ if self.multi_hop_max_dist > 0:
247
+ spatial_pos_ = ops.clip(spatial_pos_, 0, self.multi_hop_max_dist)
248
+ edge_input = edge_input[:, :, :, : self.multi_hop_max_dist, :]
249
+
250
+ # edge_input: [batch_size, num_nodes, num_nodes, max_dist, edge_feat_dim]
251
+ # edge_encoder -> [batch_size, num_nodes, num_nodes, max_dist, edge_feat_dim, num_heads]
252
+ edge_enc = ops.mean(self.edge_encoder(edge_input), axis=-2)
253
+ max_dist = ops.shape(edge_enc)[3]
254
+
255
+ # edge_enc: [batch_size, num_nodes, num_nodes, max_dist, num_heads]
256
+ # permute to [max_dist, batch_size * num_nodes * num_nodes, num_heads]
257
+ edge_enc_perm = ops.transpose(edge_enc, (3, 0, 1, 2, 4))
258
+ edge_input_flat = ops.reshape(edge_enc_perm, (max_dist, -1, self.num_heads))
259
+
260
+ # Weight shape for edge_dis_encoder: [num_edge_dis * num_heads * num_heads, 1]
261
+ if not self.edge_dis_encoder.built:
262
+ self.edge_dis_encoder.build(None)
263
+ dis_weights = ops.reshape(
264
+ self.edge_dis_encoder.weights[0], (-1, self.num_heads, self.num_heads)
265
+ )[:max_dist, :, :]
266
+
267
+ edge_input_flat = ops.matmul(edge_input_flat, dis_weights)
268
+ edge_enc_back = ops.reshape(
269
+ edge_input_flat, (max_dist, batch_size, num_nodes, num_nodes, self.num_heads)
270
+ )
271
+ # Permute to [batch_size, num_nodes, num_nodes, max_dist, num_heads]
272
+ edge_enc_back = ops.transpose(edge_enc_back, (1, 2, 3, 0, 4))
273
+
274
+ # Sum over distance dim and divide by path length:
275
+ sp_float = ops.expand_dims(ops.cast(spatial_pos_, "float32"), axis=-1)
276
+ edge_bias = ops.sum(edge_enc_back, axis=3) / sp_float
277
+ edge_bias = ops.transpose(edge_bias, (0, 3, 1, 2))
278
+ sub_bias = sub_bias + edge_bias
279
+ elif attn_edge_type is not None:
280
+ edge_bias = ops.mean(self.edge_encoder(attn_edge_type), axis=-2)
281
+ edge_bias = ops.transpose(edge_bias, (0, 3, 1, 2))
282
+ sub_bias = sub_bias + edge_bias
283
+
284
+ # Reconstruct graph_attn_bias:
285
+ top_row = ops.concatenate([corner, row0], axis=3)
286
+ bottom_rows = ops.concatenate([col0, sub_bias], axis=3)
287
+ graph_attn_bias = ops.concatenate([top_row, bottom_rows], axis=2)
288
+
289
+ # Reset padding elements with -inf mask
290
+ graph_attn_bias = graph_attn_bias + ops.expand_dims(attn_bias, axis=1)
291
+ return graph_attn_bias
292
+
293
+
294
+ class GraphormerMultiheadAttention(layers.Layer):
295
+ r"""Multi-head self-attention layer with support for additive graph structural attention bias
296
+ and key padding masks.
297
+
298
+ Args:
299
+ embed_dim (int): Total embedding dimension.
300
+ num_heads (int): Number of attention heads.
301
+ dropout (float, optional): Attention dropout probability. (default: ``0.0``)
302
+ bias (bool, optional): Whether to use bias in linear projections. (default: ``True``)
303
+ **kwargs: Additional layer arguments.
304
+
305
+ Example:
306
+ ```python
307
+ import numpy as np
308
+ from k3_node.models import GraphormerMultiheadAttention
309
+
310
+ x = np.random.rand(2, 6, 32).astype("float32") # [batch, tokens, embed_dim]
311
+ attn_bias = np.zeros((2, 4, 6, 6), dtype="float32") # structural bias per head
312
+ attn = GraphormerMultiheadAttention(embed_dim=32, num_heads=4, dropout=0.0)
313
+ out, weights = attn(x, attn_bias=attn_bias)
314
+ print(tuple(out.shape), tuple(weights.shape)) # (2, 6, 32) (2, 4, 6, 6)
315
+ ```
316
+ """
317
+
318
+ def __init__(
319
+ self,
320
+ embed_dim: int,
321
+ num_heads: int,
322
+ dropout: float = 0.0,
323
+ bias: bool = True,
324
+ **kwargs,
325
+ ):
326
+ super().__init__(**kwargs)
327
+ self.embed_dim = embed_dim
328
+ self.num_heads = num_heads
329
+ self.dropout_rate = dropout
330
+ self.head_dim = embed_dim // num_heads
331
+ if self.head_dim * num_heads != embed_dim:
332
+ raise ValueError(f"embed_dim ({embed_dim}) must be divisible by num_heads ({num_heads}).")
333
+
334
+ self.scaling = self.head_dim ** -0.5
335
+
336
+ self.q_proj = layers.Dense(embed_dim, use_bias=bias, name="q_proj")
337
+ self.k_proj = layers.Dense(embed_dim, use_bias=bias, name="k_proj")
338
+ self.v_proj = layers.Dense(embed_dim, use_bias=bias, name="v_proj")
339
+ self.out_proj = layers.Dense(embed_dim, use_bias=bias, name="out_proj")
340
+ self.dropout = layers.Dropout(dropout)
341
+
342
+ def build(self, input_shape=None):
343
+ if not self.built:
344
+ self.q_proj.build((None, None, self.embed_dim))
345
+ self.k_proj.build((None, None, self.embed_dim))
346
+ self.v_proj.build((None, None, self.embed_dim))
347
+ self.out_proj.build((None, None, self.embed_dim))
348
+ super().build(input_shape)
349
+
350
+ def call(
351
+ self,
352
+ x,
353
+ attn_bias=None,
354
+ key_padding_mask=None,
355
+ training: bool = False,
356
+ ):
357
+ r"""
358
+ Args:
359
+ x (Tensor): Sequence representation of shape ``[batch_size, seq_len, embed_dim]``.
360
+ attn_bias (Tensor, optional): Additive bias of shape
361
+ ``[batch_size, num_heads, seq_len, seq_len]``.
362
+ key_padding_mask (Tensor, optional): Boolean padding mask of shape ``[batch_size, seq_len]``,
363
+ where True indicates padding tokens to ignore.
364
+ training (bool, optional): Whether in training mode. (default: ``False``)
365
+
366
+ Returns:
367
+ Tuple[Tensor, Tensor]: Output tensor of shape ``[batch_size, seq_len, embed_dim]``
368
+ and attention weights of shape ``[batch_size, num_heads, seq_len, seq_len]``.
369
+ """
370
+ shape = ops.shape(x)
371
+ batch_size, seq_len = shape[0], shape[1]
372
+
373
+ q = self.q_proj(x) * self.scaling
374
+ k = self.k_proj(x)
375
+ v = self.v_proj(x)
376
+
377
+ # Reshape to [batch_size, num_heads, seq_len, head_dim]
378
+ q = ops.transpose(ops.reshape(q, (batch_size, seq_len, self.num_heads, self.head_dim)), (0, 2, 1, 3))
379
+ k = ops.transpose(ops.reshape(k, (batch_size, seq_len, self.num_heads, self.head_dim)), (0, 2, 1, 3))
380
+ v = ops.transpose(ops.reshape(v, (batch_size, seq_len, self.num_heads, self.head_dim)), (0, 2, 1, 3))
381
+
382
+ # [batch_size, num_heads, seq_len, seq_len]
383
+ attn_weights = ops.matmul(q, ops.transpose(k, (0, 1, 3, 2)))
384
+
385
+ if attn_bias is not None:
386
+ attn_weights = attn_weights + attn_bias
387
+
388
+ if key_padding_mask is not None:
389
+ # key_padding_mask: [batch_size, seq_len] -> [batch_size, 1, 1, seq_len]
390
+ mask = ops.expand_dims(ops.expand_dims(key_padding_mask, axis=1), axis=2)
391
+ attn_weights = ops.where(mask, -1e9, attn_weights)
392
+
393
+ attn_probs = ops.softmax(attn_weights, axis=-1)
394
+ attn_probs = self.dropout(attn_probs, training=training)
395
+
396
+ # [batch_size, num_heads, seq_len, head_dim]
397
+ attn = ops.matmul(attn_probs, v)
398
+
399
+ # Reshape to [batch_size, seq_len, embed_dim]
400
+ attn = ops.reshape(ops.transpose(attn, (0, 2, 1, 3)), (batch_size, seq_len, self.embed_dim))
401
+ out = self.out_proj(attn)
402
+ return out, attn_probs
403
+
404
+
405
+ class GraphormerGraphEncoderLayer(layers.Layer):
406
+ r"""A single Graphormer Transformer Encoder Layer, supporting Pre-LN or Post-LN,
407
+ multi-head attention with structural attention bias, and a 2-layer FFN.
408
+
409
+ Args:
410
+ embedding_dim (int, optional): Embedding dimension. (default: ``768``)
411
+ ffn_embedding_dim (int, optional): FFN intermediate dimension. (default: ``768``)
412
+ num_attention_heads (int, optional): Number of attention heads. (default: ``32``)
413
+ dropout (float, optional): Dropout probability. (default: ``0.1``)
414
+ attention_dropout (float, optional): Attention dropout. (default: ``0.1``)
415
+ activation_dropout (float, optional): FFN activation dropout. (default: ``0.1``)
416
+ activation_fn (str, optional): Activation function (``"gelu"`` or ``"relu"``). (default: ``"gelu"``)
417
+ pre_layernorm (bool, optional): Whether to use Pre-LN. (default: ``False``)
418
+ **kwargs: Additional layer arguments.
419
+
420
+ Example:
421
+ ```python
422
+ import numpy as np
423
+ from k3_node.models import GraphormerGraphEncoderLayer
424
+
425
+ x = np.random.rand(2, 6, 32).astype("float32")
426
+ attn_bias = np.zeros((2, 4, 6, 6), dtype="float32")
427
+ layer = GraphormerGraphEncoderLayer(embedding_dim=32, ffn_embedding_dim=64, num_attention_heads=4)
428
+ print(tuple(layer(x, attn_bias=attn_bias).shape)) # (2, 6, 32)
429
+ ```
430
+ """
431
+
432
+ def __init__(
433
+ self,
434
+ embedding_dim: int = 768,
435
+ ffn_embedding_dim: int = 768,
436
+ num_attention_heads: int = 32,
437
+ dropout: float = 0.1,
438
+ attention_dropout: float = 0.1,
439
+ activation_dropout: float = 0.1,
440
+ activation_fn: str = "gelu",
441
+ pre_layernorm: bool = False,
442
+ **kwargs,
443
+ ):
444
+ super().__init__(**kwargs)
445
+ self.embedding_dim = embedding_dim
446
+ self.ffn_embedding_dim = ffn_embedding_dim
447
+ self.num_attention_heads = num_attention_heads
448
+ self.dropout_rate = dropout
449
+ self.pre_layernorm = pre_layernorm
450
+
451
+ self.self_attn = GraphormerMultiheadAttention(
452
+ embed_dim=embedding_dim,
453
+ num_heads=num_attention_heads,
454
+ dropout=attention_dropout,
455
+ name="self_attn",
456
+ )
457
+ self.self_attn_layer_norm = layers.LayerNormalization(
458
+ epsilon=1e-5, name="self_attn_layer_norm"
459
+ )
460
+ self.dropout = layers.Dropout(dropout)
461
+
462
+ self.fc1 = layers.Dense(ffn_embedding_dim, name="fc1")
463
+ self.fc2 = layers.Dense(embedding_dim, name="fc2")
464
+ self.act_dropout = layers.Dropout(activation_dropout)
465
+ self.final_layer_norm = layers.LayerNormalization(
466
+ epsilon=1e-5, name="final_layer_norm"
467
+ )
468
+
469
+ if activation_fn == "gelu":
470
+ self.activation = ops.gelu
471
+ elif activation_fn == "relu":
472
+ self.activation = ops.relu
473
+ else:
474
+ self.activation = keras.activations.get(activation_fn)
475
+
476
+ def build(self, input_shape=None):
477
+ if not self.built:
478
+ self.self_attn.build((None, None, self.embedding_dim))
479
+ self.self_attn_layer_norm.build((None, None, self.embedding_dim))
480
+ self.fc1.build((None, None, self.embedding_dim))
481
+ self.fc2.build((None, None, self.ffn_embedding_dim))
482
+ self.final_layer_norm.build((None, None, self.embedding_dim))
483
+ super().build(input_shape)
484
+
485
+ def call(
486
+ self,
487
+ x,
488
+ attn_bias=None,
489
+ key_padding_mask=None,
490
+ training: bool = False,
491
+ ):
492
+ residual = x
493
+ if self.pre_layernorm:
494
+ x = self.self_attn_layer_norm(x)
495
+
496
+ x, _ = self.self_attn(
497
+ x,
498
+ attn_bias=attn_bias,
499
+ key_padding_mask=key_padding_mask,
500
+ training=training,
501
+ )
502
+ x = self.dropout(x, training=training)
503
+ x = residual + x
504
+
505
+ if not self.pre_layernorm:
506
+ x = self.self_attn_layer_norm(x)
507
+
508
+ residual = x
509
+ if self.pre_layernorm:
510
+ x = self.final_layer_norm(x)
511
+
512
+ x = self.fc2(self.act_dropout(self.activation(self.fc1(x)), training=training))
513
+ x = self.dropout(x, training=training)
514
+ x = residual + x
515
+
516
+ if not self.pre_layernorm:
517
+ x = self.final_layer_norm(x)
518
+
519
+ return x
520
+
521
+
522
+ class GraphormerGraphEncoder(layers.Layer):
523
+ r"""Graphormer Graph Encoder stack consisting of graph node features, structural attention bias,
524
+ and multiple stacked encoder layers.
525
+
526
+ Args:
527
+ num_atoms (int): Maximum number of atom types.
528
+ num_in_degree (int): Maximum in-degree value.
529
+ num_out_degree (int): Maximum out-degree value.
530
+ num_edges (int): Maximum number of edge types.
531
+ num_spatial (int): Maximum spatial distance.
532
+ num_edge_dis (int): Maximum edge distance multiplier.
533
+ edge_type (str, optional): Edge encoding mode (``"multi_hop"`` or ``"single_hop"``). (default: ``"multi_hop"``)
534
+ multi_hop_max_dist (int, optional): Max distance for multi-hop. (default: ``20``)
535
+ num_encoder_layers (int, optional): Number of encoder layers. (default: ``12``)
536
+ embedding_dim (int, optional): Embedding dimension. (default: ``768``)
537
+ ffn_embedding_dim (int, optional): FFN intermediate dimension. (default: ``768``)
538
+ num_attention_heads (int, optional): Number of attention heads. (default: ``32``)
539
+ dropout (float, optional): Dropout probability. (default: ``0.1``)
540
+ attention_dropout (float, optional): Attention dropout. (default: ``0.1``)
541
+ activation_dropout (float, optional): FFN activation dropout. (default: ``0.1``)
542
+ pre_layernorm (bool, optional): Whether to use Pre-LN. (default: ``False``)
543
+ encoder_normalize_before (bool, optional): Whether to normalize before encoder blocks. (default: ``False``)
544
+ activation_fn (str, optional): Activation function name. (default: ``"gelu"``)
545
+ **kwargs: Additional layer arguments.
546
+
547
+ Example:
548
+ ```python
549
+ import numpy as np
550
+ from k3_node.models import GraphormerGraphEncoder
551
+
552
+ # A batch of 2 graphs padded to 4 nodes, preprocessed into Graphormer's dense inputs
553
+ data = {
554
+ "x": np.random.randint(1, 15, size=(2, 4, 2)), # 2 categorical features per node
555
+ "in_degree": np.random.randint(0, 7, size=(2, 4)),
556
+ "out_degree": np.random.randint(0, 7, size=(2, 4)),
557
+ "attn_bias": np.zeros((2, 5, 5), dtype="float32"), # +1 for the virtual graph token
558
+ "spatial_pos": np.random.randint(0, 7, size=(2, 4, 4)), # shortest-path distances
559
+ "edge_input": np.random.randint(0, 7, size=(2, 4, 4, 2, 2)), # edge types along shortest paths
560
+ }
561
+
562
+ encoder = GraphormerGraphEncoder(num_atoms=16, num_in_degree=8, num_out_degree=8, num_edges=8,
563
+ num_spatial=8, num_edge_dis=4, num_encoder_layers=2,
564
+ embedding_dim=32, ffn_embedding_dim=64, num_attention_heads=4)
565
+ node_states, graph_rep = encoder(**data)
566
+ print(tuple(node_states.shape), tuple(graph_rep.shape)) # (2, 5, 32) (2, 32): per-token states, graph-token embedding
567
+ ```
568
+ """
569
+
570
+ def __init__(
571
+ self,
572
+ num_atoms: int = 512,
573
+ num_in_degree: int = 512,
574
+ num_out_degree: int = 512,
575
+ num_edges: int = 512,
576
+ num_spatial: int = 512,
577
+ num_edge_dis: int = 128,
578
+ edge_type: str = "multi_hop",
579
+ multi_hop_max_dist: int = 20,
580
+ num_encoder_layers: int = 12,
581
+ embedding_dim: int = 768,
582
+ ffn_embedding_dim: int = 768,
583
+ num_attention_heads: int = 32,
584
+ dropout: float = 0.1,
585
+ attention_dropout: float = 0.1,
586
+ activation_dropout: float = 0.1,
587
+ pre_layernorm: bool = False,
588
+ encoder_normalize_before: bool = False,
589
+ activation_fn: str = "gelu",
590
+ **kwargs,
591
+ ):
592
+ super().__init__(**kwargs)
593
+ self.embedding_dim = embedding_dim
594
+ self.num_encoder_layers = num_encoder_layers
595
+ self.pre_layernorm = pre_layernorm
596
+
597
+ self.graph_node_feature = GraphNodeFeature(
598
+ num_atoms=num_atoms,
599
+ num_in_degree=num_in_degree,
600
+ num_out_degree=num_out_degree,
601
+ hidden_dim=embedding_dim,
602
+ name="graph_node_feature",
603
+ )
604
+ self.graph_attn_bias = GraphAttnBias(
605
+ num_heads=num_attention_heads,
606
+ num_atoms=num_atoms,
607
+ num_edges=num_edges,
608
+ num_spatial=num_spatial,
609
+ num_edge_dis=num_edge_dis,
610
+ edge_type=edge_type,
611
+ multi_hop_max_dist=multi_hop_max_dist,
612
+ name="graph_attn_bias",
613
+ )
614
+ if encoder_normalize_before:
615
+ self.emb_layer_norm = layers.LayerNormalization(epsilon=1e-5, name="emb_layer_norm")
616
+ else:
617
+ self.emb_layer_norm = None
618
+
619
+ self.dropout = layers.Dropout(dropout)
620
+
621
+ self.encoder_layers = [
622
+ GraphormerGraphEncoderLayer(
623
+ embedding_dim=embedding_dim,
624
+ ffn_embedding_dim=ffn_embedding_dim,
625
+ num_attention_heads=num_attention_heads,
626
+ dropout=dropout,
627
+ attention_dropout=attention_dropout,
628
+ activation_dropout=activation_dropout,
629
+ activation_fn=activation_fn,
630
+ pre_layernorm=pre_layernorm,
631
+ name=f"layers_{i}",
632
+ )
633
+ for i in range(num_encoder_layers)
634
+ ]
635
+
636
+ if pre_layernorm:
637
+ self.final_layer_norm = layers.LayerNormalization(epsilon=1e-5, name="final_layer_norm")
638
+ else:
639
+ self.final_layer_norm = None
640
+
641
+ def build(self, input_shape=None):
642
+ if not self.built:
643
+ self.graph_node_feature.build(None)
644
+ self.graph_attn_bias.build(None)
645
+ if self.emb_layer_norm is not None:
646
+ self.emb_layer_norm.build((None, None, self.embedding_dim))
647
+ for layer in self.encoder_layers:
648
+ layer.build((None, None, self.embedding_dim))
649
+ if self.final_layer_norm is not None:
650
+ self.final_layer_norm.build((None, None, self.embedding_dim))
651
+ super().build(input_shape)
652
+
653
+ def call(
654
+ self,
655
+ x,
656
+ in_degree,
657
+ out_degree,
658
+ attn_bias,
659
+ spatial_pos,
660
+ edge_input=None,
661
+ attn_edge_type=None,
662
+ perturb=None,
663
+ training: bool = False,
664
+ ):
665
+ batch_size = ops.shape(x)[0]
666
+ # Compute padding mask: [batch_size, num_nodes]
667
+ if len(ops.shape(x)) == 3:
668
+ raw_mask = ops.equal(x[:, :, 0], 0)
669
+ else:
670
+ raw_mask = ops.equal(x, 0)
671
+
672
+ # Prepend False for graph token: [batch_size, num_nodes + 1]
673
+ cls_mask = ops.zeros((batch_size, 1), dtype="bool")
674
+ padding_mask = ops.concatenate([cls_mask, raw_mask], axis=1)
675
+
676
+ # Node features: [batch_size, num_nodes + 1, embedding_dim]
677
+ h = self.graph_node_feature(x, in_degree, out_degree)
678
+ if perturb is not None:
679
+ # perturb is added to non-token nodes
680
+ h_token = h[:, 0:1, :]
681
+ h_nodes = h[:, 1:, :] + perturb
682
+ h = ops.concatenate([h_token, h_nodes], axis=1)
683
+
684
+ bias = self.graph_attn_bias(
685
+ attn_bias=attn_bias,
686
+ spatial_pos=spatial_pos,
687
+ x=x,
688
+ edge_input=edge_input,
689
+ attn_edge_type=attn_edge_type,
690
+ )
691
+
692
+ if self.emb_layer_norm is not None:
693
+ h = self.emb_layer_norm(h)
694
+
695
+ h = self.dropout(h, training=training)
696
+
697
+ for layer in self.encoder_layers:
698
+ h = layer(
699
+ h,
700
+ attn_bias=bias,
701
+ key_padding_mask=padding_mask,
702
+ training=training,
703
+ )
704
+
705
+ if self.final_layer_norm is not None:
706
+ h = self.final_layer_norm(h)
707
+
708
+ graph_rep = h[:, 0, :]
709
+ return h, graph_rep
710
+
711
+
712
+ class Graphormer(keras.Model):
713
+ r"""Graphormer model for molecular graph representation and property prediction
714
+ from `"Do Transformers Really Perform Badly for Graph Representation?" <https://arxiv.org/abs/2106.05234>`_.
715
+
716
+ Args:
717
+ num_atoms (int, optional): Maximum atom vocabulary size. (default: ``512``)
718
+ num_in_degree (int, optional): Maximum in-degree. (default: ``512``)
719
+ num_out_degree (int, optional): Maximum out-degree. (default: ``512``)
720
+ num_edges (int, optional): Maximum edge vocabulary size. (default: ``512``)
721
+ num_spatial (int, optional): Maximum spatial distance. (default: ``512``)
722
+ num_edge_dis (int, optional): Maximum edge distance multiplier. (default: ``128``)
723
+ edge_type (str, optional): Type of edge encoding (``"multi_hop"`` or ``"single_hop"``). (default: ``"multi_hop"``)
724
+ multi_hop_max_dist (int, optional): Max distance for multi-hop paths. (default: ``20``)
725
+ num_encoder_layers (int, optional): Number of Transformer layers. (default: ``12``)
726
+ embedding_dim (int, optional): Hidden embedding dimension. (default: ``768``)
727
+ ffn_embedding_dim (int, optional): FFN intermediate dimension. (default: ``768``)
728
+ num_attention_heads (int, optional): Number of attention heads. (default: ``32``)
729
+ dropout (float, optional): Dropout rate. (default: ``0.0``)
730
+ attention_dropout (float, optional): Attention dropout rate. (default: ``0.1``)
731
+ activation_dropout (float, optional): Activation dropout rate. (default: ``0.1``)
732
+ encoder_normalize_before (bool, optional): Whether to apply LayerNorm before encoder blocks. (default: ``True``)
733
+ pre_layernorm (bool, optional): Whether to use Pre-LN. (default: ``False``)
734
+ num_classes (int, optional): Output dimension for prediction head. (default: ``1``)
735
+ activation_fn (str, optional): Activation function. (default: ``"gelu"``)
736
+ **kwargs: Additional model arguments.
737
+
738
+ Example:
739
+ ```python
740
+ import numpy as np
741
+ from k3_node.models import Graphormer
742
+
743
+ # A batch of 2 graphs padded to 4 nodes, preprocessed into Graphormer's dense inputs
744
+ data = {
745
+ "x": np.random.randint(1, 15, size=(2, 4, 2)), # 2 categorical features per node
746
+ "in_degree": np.random.randint(0, 7, size=(2, 4)),
747
+ "out_degree": np.random.randint(0, 7, size=(2, 4)),
748
+ "attn_bias": np.zeros((2, 5, 5), dtype="float32"), # +1 for the virtual graph token
749
+ "spatial_pos": np.random.randint(0, 7, size=(2, 4, 4)), # shortest-path distances
750
+ "edge_input": np.random.randint(0, 7, size=(2, 4, 4, 2, 2)), # edge types along shortest paths
751
+ }
752
+
753
+ model = Graphormer(num_atoms=16, num_in_degree=8, num_out_degree=8, num_edges=8, num_spatial=8,
754
+ num_edge_dis=4, num_encoder_layers=2, embedding_dim=32, ffn_embedding_dim=64,
755
+ num_attention_heads=4, num_classes=1)
756
+ out = model(data) # one prediction per graph
757
+ print(tuple(out.shape)) # (2, 1)
758
+ ```
759
+ """
760
+
761
+ def __init__(
762
+ self,
763
+ num_atoms: int = 512,
764
+ num_in_degree: int = 512,
765
+ num_out_degree: int = 512,
766
+ num_edges: int = 512,
767
+ num_spatial: int = 512,
768
+ num_edge_dis: int = 128,
769
+ edge_type: str = "multi_hop",
770
+ multi_hop_max_dist: int = 20,
771
+ num_encoder_layers: int = 12,
772
+ embedding_dim: int = 768,
773
+ ffn_embedding_dim: int = 768,
774
+ num_attention_heads: int = 32,
775
+ dropout: float = 0.0,
776
+ attention_dropout: float = 0.1,
777
+ activation_dropout: float = 0.1,
778
+ encoder_normalize_before: bool = True,
779
+ pre_layernorm: bool = False,
780
+ num_classes: int = 1,
781
+ activation_fn: str = "gelu",
782
+ **kwargs,
783
+ ):
784
+ super().__init__(**kwargs)
785
+ self.num_atoms = num_atoms
786
+ self.num_in_degree = num_in_degree
787
+ self.num_out_degree = num_out_degree
788
+ self.num_edges = num_edges
789
+ self.num_spatial = num_spatial
790
+ self.num_edge_dis = num_edge_dis
791
+ self.edge_type = edge_type
792
+ self.multi_hop_max_dist = multi_hop_max_dist
793
+ self.num_encoder_layers = num_encoder_layers
794
+ self.embedding_dim = embedding_dim
795
+ self.ffn_embedding_dim = ffn_embedding_dim
796
+ self.num_attention_heads = num_attention_heads
797
+ self.num_classes = num_classes
798
+ self.pre_layernorm = pre_layernorm
799
+
800
+ self.graph_encoder = GraphormerGraphEncoder(
801
+ num_atoms=num_atoms,
802
+ num_in_degree=num_in_degree,
803
+ num_out_degree=num_out_degree,
804
+ num_edges=num_edges,
805
+ num_spatial=num_spatial,
806
+ num_edge_dis=num_edge_dis,
807
+ edge_type=edge_type,
808
+ multi_hop_max_dist=multi_hop_max_dist,
809
+ num_encoder_layers=num_encoder_layers,
810
+ embedding_dim=embedding_dim,
811
+ ffn_embedding_dim=ffn_embedding_dim,
812
+ num_attention_heads=num_attention_heads,
813
+ dropout=dropout,
814
+ attention_dropout=attention_dropout,
815
+ activation_dropout=activation_dropout,
816
+ encoder_normalize_before=encoder_normalize_before,
817
+ pre_layernorm=pre_layernorm,
818
+ activation_fn=activation_fn,
819
+ name="graph_encoder",
820
+ )
821
+
822
+ # Output prediction head (matching GraphormerEncoder in Graphormer-main)
823
+ self.lm_head_transform_weight = layers.Dense(
824
+ embedding_dim, name="lm_head_transform_weight"
825
+ )
826
+ self.layer_norm = layers.LayerNormalization(epsilon=1e-5, name="layer_norm")
827
+ self.embed_out = layers.Dense(num_classes, use_bias=False, name="embed_out")
828
+
829
+ if activation_fn == "gelu":
830
+ self.head_act = ops.gelu
831
+ elif activation_fn == "relu":
832
+ self.head_act = ops.relu
833
+ else:
834
+ self.head_act = keras.activations.get(activation_fn)
835
+
836
+ def build(self, input_shape=None):
837
+ if not self.built:
838
+ self.graph_encoder.build(None)
839
+ self.lm_head_transform_weight.build((None, self.embedding_dim))
840
+ self.layer_norm.build((None, self.embedding_dim))
841
+ self.embed_out.build((None, self.embedding_dim))
842
+ self.lm_output_learned_bias = self.add_weight(
843
+ shape=(1,),
844
+ initializer="zeros",
845
+ trainable=True,
846
+ name="lm_output_learned_bias",
847
+ )
848
+ super().build(input_shape)
849
+
850
+ def call(
851
+ self,
852
+ batched_data: Optional[Dict[str, Any]] = None,
853
+ x=None,
854
+ in_degree=None,
855
+ out_degree=None,
856
+ attn_bias=None,
857
+ spatial_pos=None,
858
+ edge_input=None,
859
+ attn_edge_type=None,
860
+ perturb=None,
861
+ return_all: bool = False,
862
+ training: bool = False,
863
+ ):
864
+ r"""Forward pass for Graphormer.
865
+
866
+ Accepts either a single dictionary ``batched_data`` containing graph tensors,
867
+ or individual tensor arguments.
868
+
869
+ Returns:
870
+ Tensor or Tuple[Tensor, Tensor]: Graph prediction tensor of shape ``[batch_size, num_classes]``,
871
+ or if return_all=True, a tuple of ``(graph_pred, all_node_features)``.
872
+ """
873
+ if batched_data is not None:
874
+ x = batched_data["x"]
875
+ in_degree = batched_data["in_degree"]
876
+ out_degree = batched_data["out_degree"]
877
+ attn_bias = batched_data["attn_bias"]
878
+ spatial_pos = batched_data["spatial_pos"]
879
+ edge_input = batched_data.get("edge_input", None)
880
+ attn_edge_type = batched_data.get("attn_edge_type", None)
881
+
882
+ h, graph_rep = self.graph_encoder(
883
+ x=x,
884
+ in_degree=in_degree,
885
+ out_degree=out_degree,
886
+ attn_bias=attn_bias,
887
+ spatial_pos=spatial_pos,
888
+ edge_input=edge_input,
889
+ attn_edge_type=attn_edge_type,
890
+ perturb=perturb,
891
+ training=training,
892
+ )
893
+
894
+ # Output projection on graph token: shape [batch_size, embedding_dim]
895
+ token_h = self.layer_norm(self.head_act(self.lm_head_transform_weight(graph_rep)))
896
+ out = self.embed_out(token_h)
897
+ if hasattr(self, "lm_output_learned_bias") and self.lm_output_learned_bias is not None:
898
+ out = out + self.lm_output_learned_bias
899
+
900
+ if return_all:
901
+ return out, h
902
+ return out
903
+
904
+ def embed(
905
+ self,
906
+ batched_data: Optional[Dict[str, Any]] = None,
907
+ x=None,
908
+ in_degree=None,
909
+ out_degree=None,
910
+ attn_bias=None,
911
+ spatial_pos=None,
912
+ edge_input=None,
913
+ attn_edge_type=None,
914
+ ):
915
+ r"""Computes graph and node embeddings without applying the prediction head."""
916
+ if batched_data is not None:
917
+ x = batched_data["x"]
918
+ in_degree = batched_data["in_degree"]
919
+ out_degree = batched_data["out_degree"]
920
+ attn_bias = batched_data["attn_bias"]
921
+ spatial_pos = batched_data["spatial_pos"]
922
+ edge_input = batched_data.get("edge_input", None)
923
+ attn_edge_type = batched_data.get("attn_edge_type", None)
924
+
925
+ h, graph_rep = self.graph_encoder(
926
+ x=x,
927
+ in_degree=in_degree,
928
+ out_degree=out_degree,
929
+ attn_bias=attn_bias,
930
+ spatial_pos=spatial_pos,
931
+ edge_input=edge_input,
932
+ attn_edge_type=attn_edge_type,
933
+ training=False,
934
+ )
935
+ return graph_rep
936
+
937
+ @classmethod
938
+ def from_pretrained(
939
+ cls,
940
+ pretrained_name: str = "pcqm4mv1_graphormer_base",
941
+ folder: str = "checkpoints",
942
+ download: bool = True,
943
+ **kwargs,
944
+ ) -> "Graphormer":
945
+ r"""Instantiates a Graphormer model with pre-trained weights."""
946
+ cfg = get_graphormer_config(pretrained_name)
947
+ cfg.update(kwargs)
948
+ model = cls(**cfg)
949
+ load_graphormer_weights(model, pretrained_name=pretrained_name, folder=folder, download=download)
950
+ return model
951
+
952
+ def __repr__(self) -> str:
953
+ return (
954
+ f"{self.__class__.__name__}("
955
+ f"embedding_dim={self.embedding_dim}, "
956
+ f"num_encoder_layers={self.num_encoder_layers}, "
957
+ f"num_attention_heads={self.num_attention_heads}, "
958
+ f"num_classes={self.num_classes})"
959
+ )
960
+
961
+
962
+ # =========================================================================
963
+ # Configuration presets & Pre-trained URLs
964
+ # =========================================================================
965
+
966
+ PRETRAINED_MODEL_URLS = {
967
+ "pcqm4mv1_graphormer_base": "https://huggingface.co/clefourrier/graphormer-base-pcqm4mv1/resolve/main/pytorch_model.bin",
968
+ "pcqm4mv2_graphormer_base": "https://huggingface.co/clefourrier/graphormer-base-pcqm4mv2/resolve/main/pytorch_model.bin",
969
+ "pcqm4mv1_graphormer_base_for_molhiv": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_base_preln_pcqm4mv1_for_hiv.pt",
970
+ }
971
+
972
+ LEGACY_URLS = {
973
+ "pcqm4mv1_graphormer_base": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv1.pt",
974
+ "pcqm4mv2_graphormer_base": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv2.pt",
975
+ }
976
+
977
+
978
+ def get_graphormer_config(name_or_variant: str) -> Dict[str, Any]:
979
+ r"""Returns configuration dictionary for known Graphormer architectures."""
980
+ key = name_or_variant.lower().replace("-", "_")
981
+ if "slim" in key:
982
+ return {
983
+ "num_encoder_layers": 12,
984
+ "num_attention_heads": 8,
985
+ "embedding_dim": 80,
986
+ "ffn_embedding_dim": 80,
987
+ "pre_layernorm": False,
988
+ "encoder_normalize_before": True,
989
+ "dropout": 0.0,
990
+ "attention_dropout": 0.1,
991
+ "activation_dropout": 0.1,
992
+ }
993
+ elif "large" in key:
994
+ return {
995
+ "num_encoder_layers": 24,
996
+ "num_attention_heads": 32,
997
+ "embedding_dim": 1024,
998
+ "ffn_embedding_dim": 1024,
999
+ "pre_layernorm": False,
1000
+ "encoder_normalize_before": True,
1001
+ "dropout": 0.0,
1002
+ "attention_dropout": 0.1,
1003
+ "activation_dropout": 0.1,
1004
+ }
1005
+ elif "base" in key or "pcqm4m" in key:
1006
+ pre_ln = "molhiv" in key or "preln" in key
1007
+ num_classes = 1
1008
+ if "molhiv" in key:
1009
+ num_classes = 2
1010
+ return {
1011
+ "num_encoder_layers": 12,
1012
+ "num_attention_heads": 32,
1013
+ "embedding_dim": 768,
1014
+ "ffn_embedding_dim": 768,
1015
+ "pre_layernorm": pre_ln,
1016
+ "encoder_normalize_before": True,
1017
+ "dropout": 0.0,
1018
+ "attention_dropout": 0.1,
1019
+ "activation_dropout": 0.1,
1020
+ "num_classes": num_classes,
1021
+ }
1022
+ else:
1023
+ # Standard default architecture
1024
+ return {
1025
+ "num_encoder_layers": 6,
1026
+ "num_attention_heads": 8,
1027
+ "embedding_dim": 1024,
1028
+ "ffn_embedding_dim": 4096,
1029
+ "pre_layernorm": False,
1030
+ "encoder_normalize_before": True,
1031
+ "dropout": 0.1,
1032
+ "attention_dropout": 0.1,
1033
+ "activation_dropout": 0.0,
1034
+ "num_classes": 1,
1035
+ }
1036
+
1037
+
1038
+ def download_graphormer_checkpoint(
1039
+ name: str,
1040
+ folder: str = "checkpoints",
1041
+ log: bool = True,
1042
+ ) -> str:
1043
+ r"""Downloads a pre-trained Graphormer checkpoint to the specified folder.
1044
+
1045
+ Args:
1046
+ name (str): Pretrained checkpoint name (e.g., ``"pcqm4mv1_graphormer_base"``,
1047
+ ``"pcqm4mv2_graphormer_base"``).
1048
+ folder (str, optional): Destination folder. (default: ``"checkpoints"``)
1049
+ log (bool, optional): Whether to print download progress. (default: ``True``)
1050
+
1051
+ Returns:
1052
+ str: Absolute path to the downloaded file.
1053
+ """
1054
+ from k3_node.data.download import download_url
1055
+
1056
+ clean_name = name.lower().replace("-", "_")
1057
+ if clean_name not in PRETRAINED_MODEL_URLS and name not in PRETRAINED_MODEL_URLS:
1058
+ raise ValueError(
1059
+ f"Unknown pretrained model name '{name}'. Available: {list(PRETRAINED_MODEL_URLS.keys())}"
1060
+ )
1061
+
1062
+ matched_key = clean_name if clean_name in PRETRAINED_MODEL_URLS else name
1063
+ url = PRETRAINED_MODEL_URLS[matched_key]
1064
+ filename = f"{matched_key}.pt"
1065
+ local_path = os.path.join(folder, filename)
1066
+
1067
+ if os.path.exists(local_path):
1068
+ return local_path
1069
+
1070
+ # Check local repo path if present
1071
+ repo_alt = os.path.join("Graphormer-main", "checkpoints", filename)
1072
+ if os.path.exists(repo_alt):
1073
+ return repo_alt
1074
+
1075
+ try:
1076
+ return download_url(url, folder=folder, filename=filename, log=log)
1077
+ except Exception as e:
1078
+ if matched_key in LEGACY_URLS:
1079
+ fallback_url = LEGACY_URLS[matched_key]
1080
+ try:
1081
+ return download_url(fallback_url, folder=folder, filename=filename, log=log)
1082
+ except Exception:
1083
+ pass
1084
+ raise RuntimeError(f"Failed to download checkpoint for '{name}' from {url}: {e}")
1085
+
1086
+
1087
+ def load_graphormer_weights(
1088
+ model: Graphormer,
1089
+ checkpoint_path: Optional[str] = None,
1090
+ pretrained_name: Optional[str] = None,
1091
+ folder: str = "checkpoints",
1092
+ download: bool = True,
1093
+ ) -> Graphormer:
1094
+ r"""Loads weights from a PyTorch (.pt or .bin) checkpoint into a Keras Graphormer model.
1095
+
1096
+ Args:
1097
+ model (Graphormer): Target model instance.
1098
+ checkpoint_path (str, optional): Path to .pt or .bin file.
1099
+ pretrained_name (str, optional): Name of pre-trained model to load or download.
1100
+ folder (str, optional): Checkpoints directory. (default: ``"checkpoints"``)
1101
+ download (bool, optional): Whether to download if missing. (default: ``True``)
1102
+
1103
+ Returns:
1104
+ Graphormer: The model with loaded weights.
1105
+ """
1106
+ path_to_load = checkpoint_path
1107
+
1108
+ if path_to_load is None:
1109
+ if pretrained_name is None:
1110
+ raise ValueError("Either checkpoint_path or pretrained_name must be specified.")
1111
+ candidate = os.path.join(folder, f"{pretrained_name}.pt")
1112
+ if os.path.isfile(candidate):
1113
+ path_to_load = candidate
1114
+ elif download:
1115
+ path_to_load = download_graphormer_checkpoint(pretrained_name, folder=folder)
1116
+ else:
1117
+ raise FileNotFoundError(f"Checkpoint for '{pretrained_name}' not found at '{candidate}'.")
1118
+ elif not os.path.isfile(path_to_load) and download and pretrained_name:
1119
+ path_to_load = download_graphormer_checkpoint(pretrained_name, folder=folder)
1120
+
1121
+ import torch
1122
+ import numpy as np
1123
+
1124
+ state = torch.load(path_to_load, map_location="cpu")
1125
+ if isinstance(state, dict) and "model" in state:
1126
+ state_dict = state["model"]
1127
+ elif isinstance(state, dict):
1128
+ state_dict = state
1129
+ else:
1130
+ raise ValueError(f"Unexpected checkpoint state format: {type(state)}")
1131
+
1132
+ if not model.built:
1133
+ model.build(None)
1134
+
1135
+ def _to_tensor(t):
1136
+ if hasattr(t, "detach"):
1137
+ t = t.detach()
1138
+ if hasattr(t, "numpy"):
1139
+ t = t.numpy()
1140
+ return ops.convert_to_tensor(np.array(t, dtype=np.float32), dtype="float32")
1141
+
1142
+ # Prefix stripping (fairseq checkpoints often have 'encoder.' or 'graph_encoder.')
1143
+ clean_dict = {}
1144
+ for k, v in state_dict.items():
1145
+ ck = k
1146
+ if ck.startswith("encoder."):
1147
+ ck = ck[len("encoder.") :]
1148
+ clean_dict[ck] = v
1149
+
1150
+ # 1. GraphNodeFeature
1151
+ gnf = model.graph_encoder.graph_node_feature
1152
+ for param_name, layer in [
1153
+ ("atom_encoder.weight", gnf.atom_encoder),
1154
+ ("in_degree_encoder.weight", gnf.in_degree_encoder),
1155
+ ("out_degree_encoder.weight", gnf.out_degree_encoder),
1156
+ ("graph_token.weight", gnf.graph_token),
1157
+ ]:
1158
+ key = f"graph_encoder.graph_node_feature.{param_name}"
1159
+ if key in clean_dict:
1160
+ layer.weights[0].assign(_to_tensor(clean_dict[key]))
1161
+
1162
+ # 2. GraphAttnBias
1163
+ gab = model.graph_encoder.graph_attn_bias
1164
+ for param_name, layer in [
1165
+ ("edge_encoder.weight", gab.edge_encoder),
1166
+ ("spatial_pos_encoder.weight", gab.spatial_pos_encoder),
1167
+ ("graph_token_virtual_distance.weight", gab.graph_token_virtual_distance),
1168
+ ]:
1169
+ key = f"graph_encoder.graph_attn_bias.{param_name}"
1170
+ if key in clean_dict:
1171
+ layer.weights[0].assign(_to_tensor(clean_dict[key]))
1172
+
1173
+ if hasattr(gab, "edge_dis_encoder"):
1174
+ key = "graph_encoder.graph_attn_bias.edge_dis_encoder.weight"
1175
+ if key in clean_dict:
1176
+ gab.edge_dis_encoder.weights[0].assign(_to_tensor(clean_dict[key]))
1177
+
1178
+ # 3. emb_layer_norm
1179
+ if model.graph_encoder.emb_layer_norm is not None:
1180
+ eln = model.graph_encoder.emb_layer_norm
1181
+ if "graph_encoder.emb_layer_norm.weight" in clean_dict:
1182
+ eln.gamma.assign(_to_tensor(clean_dict["graph_encoder.emb_layer_norm.weight"]))
1183
+ if "graph_encoder.emb_layer_norm.bias" in clean_dict:
1184
+ eln.beta.assign(_to_tensor(clean_dict["graph_encoder.emb_layer_norm.bias"]))
1185
+
1186
+ # 4. Encoder layers
1187
+ for i, enc_layer in enumerate(model.graph_encoder.encoder_layers):
1188
+ p = f"graph_encoder.layers.{i}"
1189
+ # Attention projections
1190
+ for proj_name in ["q_proj", "k_proj", "v_proj", "out_proj"]:
1191
+ proj = getattr(enc_layer.self_attn, proj_name)
1192
+ w_key = f"{p}.self_attn.{proj_name}.weight"
1193
+ b_key = f"{p}.self_attn.{proj_name}.bias"
1194
+ if w_key in clean_dict:
1195
+ proj.kernel.assign(_to_tensor(clean_dict[w_key].t()))
1196
+ if b_key in clean_dict and proj.bias is not None:
1197
+ proj.bias.assign(_to_tensor(clean_dict[b_key]))
1198
+
1199
+ # self_attn_layer_norm
1200
+ if f"{p}.self_attn_layer_norm.weight" in clean_dict:
1201
+ enc_layer.self_attn_layer_norm.gamma.assign(
1202
+ _to_tensor(clean_dict[f"{p}.self_attn_layer_norm.weight"])
1203
+ )
1204
+ if f"{p}.self_attn_layer_norm.bias" in clean_dict:
1205
+ enc_layer.self_attn_layer_norm.beta.assign(
1206
+ _to_tensor(clean_dict[f"{p}.self_attn_layer_norm.bias"])
1207
+ )
1208
+
1209
+ # fc1 & fc2
1210
+ if f"{p}.fc1.weight" in clean_dict:
1211
+ enc_layer.fc1.kernel.assign(_to_tensor(clean_dict[f"{p}.fc1.weight"].t()))
1212
+ if f"{p}.fc1.bias" in clean_dict:
1213
+ enc_layer.fc1.bias.assign(_to_tensor(clean_dict[f"{p}.fc1.bias"]))
1214
+ if f"{p}.fc2.weight" in clean_dict:
1215
+ enc_layer.fc2.kernel.assign(_to_tensor(clean_dict[f"{p}.fc2.weight"].t()))
1216
+ if f"{p}.fc2.bias" in clean_dict:
1217
+ enc_layer.fc2.bias.assign(_to_tensor(clean_dict[f"{p}.fc2.bias"]))
1218
+
1219
+ # final_layer_norm
1220
+ if f"{p}.final_layer_norm.weight" in clean_dict:
1221
+ enc_layer.final_layer_norm.gamma.assign(
1222
+ _to_tensor(clean_dict[f"{p}.final_layer_norm.weight"])
1223
+ )
1224
+ if f"{p}.final_layer_norm.bias" in clean_dict:
1225
+ enc_layer.final_layer_norm.beta.assign(
1226
+ _to_tensor(clean_dict[f"{p}.final_layer_norm.bias"])
1227
+ )
1228
+
1229
+ # 5. final_layer_norm of encoder stack (if pre_layernorm)
1230
+ if model.graph_encoder.final_layer_norm is not None:
1231
+ fln = model.graph_encoder.final_layer_norm
1232
+ if "graph_encoder.final_layer_norm.weight" in clean_dict:
1233
+ fln.gamma.assign(_to_tensor(clean_dict["graph_encoder.final_layer_norm.weight"]))
1234
+ if "graph_encoder.final_layer_norm.bias" in clean_dict:
1235
+ fln.beta.assign(_to_tensor(clean_dict["graph_encoder.final_layer_norm.bias"]))
1236
+
1237
+ # 6. Prediction Head
1238
+ if "lm_head_transform_weight.weight" in clean_dict:
1239
+ model.lm_head_transform_weight.kernel.assign(
1240
+ _to_tensor(clean_dict["lm_head_transform_weight.weight"].t())
1241
+ )
1242
+ if "lm_head_transform_weight.bias" in clean_dict:
1243
+ model.lm_head_transform_weight.bias.assign(
1244
+ _to_tensor(clean_dict["lm_head_transform_weight.bias"])
1245
+ )
1246
+
1247
+ if "layer_norm.weight" in clean_dict:
1248
+ model.layer_norm.gamma.assign(_to_tensor(clean_dict["layer_norm.weight"]))
1249
+ if "layer_norm.bias" in clean_dict:
1250
+ model.layer_norm.beta.assign(_to_tensor(clean_dict["layer_norm.bias"]))
1251
+
1252
+ if "embed_out.weight" in clean_dict and hasattr(model, "embed_out"):
1253
+ model.embed_out.kernel.assign(_to_tensor(clean_dict["embed_out.weight"].t()))
1254
+
1255
+ if "lm_output_learned_bias" in clean_dict and hasattr(model, "lm_output_learned_bias"):
1256
+ model.lm_output_learned_bias.assign(_to_tensor(clean_dict["lm_output_learned_bias"]))
1257
+
1258
+ return model