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,1122 @@
1
+ import os
2
+ import math
3
+ import zipfile
4
+ import urllib.request
5
+ from typing import Optional, List, Union, Tuple, Dict, Any
6
+
7
+ import numpy as np
8
+ import keras
9
+ from keras import layers, ops
10
+
11
+ from k3_node.layers.pool import global_add_pool, global_max_pool, global_mean_pool
12
+ from k3_node.ops.segment import segment_sum
13
+
14
+
15
+ def _get_act(act_str: str):
16
+ act_str = act_str.lower()
17
+ if act_str in ["gelu", "quick_gelu"]:
18
+ return ops.gelu
19
+ elif act_str == "relu":
20
+ return ops.relu
21
+ elif act_str == "silu" or act_str == "swish":
22
+ return ops.silu
23
+ elif act_str == "tanh":
24
+ return ops.tanh
25
+ elif act_str == "sigmoid":
26
+ return ops.sigmoid
27
+ return ops.relu
28
+
29
+
30
+ # 2022 OGB molecule atom feature dimensions (used in GraphGPS pretraining on PCQM4Mv2)
31
+ DEFAULT_ATOM_FEATURE_DIMS = [119, 4, 12, 12, 10, 6, 6, 2, 2]
32
+ DEFAULT_BOND_FEATURE_DIMS = [5, 6, 2]
33
+
34
+
35
+ class AtomEncoder(layers.Layer):
36
+ r"""OGB Molecule categorical atom feature encoder.
37
+
38
+ Args:
39
+ emb_dim (int): Output embedding dimension.
40
+ feature_dims (List[int], optional): Categorical feature vocabulary sizes for each
41
+ atom feature column. (default: ``[119, 4, 12, 12, 10, 6, 6, 2, 2]``)
42
+ **kwargs: Additional layer arguments.
43
+
44
+ Example:
45
+ ```python
46
+ import numpy as np
47
+ from k3_node.models import AtomEncoder
48
+
49
+ # OGB-style integer atom and bond features (9 per atom, 3 per bond)
50
+ x = np.random.randint(0, 2, size=(5, 9))
51
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
52
+ edge_attr = np.random.randint(0, 2, size=(6, 3))
53
+ batch = np.array([0, 0, 0, 1, 1]) # two graphs
54
+
55
+ print(tuple(AtomEncoder(emb_dim=32)(x).shape)) # (5, 32): sum of per-feature embeddings
56
+ ```
57
+ """
58
+
59
+ def __init__(
60
+ self,
61
+ emb_dim: int,
62
+ feature_dims: Optional[List[int]] = None,
63
+ **kwargs,
64
+ ):
65
+ super().__init__(**kwargs)
66
+ self.emb_dim = emb_dim
67
+ self.feature_dims = feature_dims if feature_dims is not None else DEFAULT_ATOM_FEATURE_DIMS
68
+
69
+ self.atom_embedding_list = [
70
+ layers.Embedding(
71
+ input_dim=dim,
72
+ output_dim=emb_dim,
73
+ embeddings_initializer="glorot_uniform",
74
+ name=f"atom_embedding_{i}",
75
+ )
76
+ for i, dim in enumerate(self.feature_dims)
77
+ ]
78
+
79
+ def build(self, input_shape=None):
80
+ for emb in self.atom_embedding_list:
81
+ if not emb.built:
82
+ emb.build(None)
83
+ super().build(input_shape)
84
+
85
+ def call(self, x):
86
+ x = ops.cast(x, "int32")
87
+ out = 0
88
+ for i, emb in enumerate(self.atom_embedding_list):
89
+ feat = x[:, i]
90
+ out = out + emb(feat)
91
+ return out
92
+
93
+
94
+ class BondEncoder(layers.Layer):
95
+ r"""OGB Molecule categorical bond feature encoder.
96
+
97
+ Args:
98
+ emb_dim (int): Output embedding dimension.
99
+ feature_dims (List[int], optional): Categorical feature vocabulary sizes for each
100
+ bond feature column. (default: ``[5, 6, 2]``)
101
+ **kwargs: Additional layer arguments.
102
+
103
+ Example:
104
+ ```python
105
+ import numpy as np
106
+ from k3_node.models import BondEncoder
107
+
108
+ # OGB-style integer atom and bond features (9 per atom, 3 per bond)
109
+ x = np.random.randint(0, 2, size=(5, 9))
110
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
111
+ edge_attr = np.random.randint(0, 2, size=(6, 3))
112
+ batch = np.array([0, 0, 0, 1, 1]) # two graphs
113
+
114
+ print(tuple(BondEncoder(emb_dim=32)(edge_attr).shape)) # (6, 32)
115
+ ```
116
+ """
117
+
118
+ def __init__(
119
+ self,
120
+ emb_dim: int,
121
+ feature_dims: Optional[List[int]] = None,
122
+ **kwargs,
123
+ ):
124
+ super().__init__(**kwargs)
125
+ self.emb_dim = emb_dim
126
+ self.feature_dims = feature_dims if feature_dims is not None else DEFAULT_BOND_FEATURE_DIMS
127
+
128
+ self.bond_embedding_list = [
129
+ layers.Embedding(
130
+ input_dim=dim,
131
+ output_dim=emb_dim,
132
+ embeddings_initializer="glorot_uniform",
133
+ name=f"bond_embedding_{i}",
134
+ )
135
+ for i, dim in enumerate(self.feature_dims)
136
+ ]
137
+
138
+ def build(self, input_shape=None):
139
+ for emb in self.bond_embedding_list:
140
+ if not emb.built:
141
+ emb.build(None)
142
+ super().build(input_shape)
143
+
144
+ def call(self, edge_attr):
145
+ edge_attr = ops.cast(edge_attr, "int32")
146
+ out = 0
147
+ for i, emb in enumerate(self.bond_embedding_list):
148
+ feat = edge_attr[:, i]
149
+ out = out + emb(feat)
150
+ return out
151
+
152
+
153
+ class RWSEEncoder(layers.Layer):
154
+ r"""Random Walk Structural Encoding (RWSE) node encoder.
155
+
156
+ Normalizes the precomputed $k$-step diagonal random walk landing probabilities
157
+ using Batch Normalization and projects them into `pe_dim` dimension.
158
+
159
+ Args:
160
+ num_rw_steps (int): Number of random walk steps. (default: ``16``)
161
+ pe_dim (int): Output structural encoding dimension. (default: ``20``)
162
+ **kwargs: Additional layer arguments.
163
+
164
+ Example:
165
+ ```python
166
+ import numpy as np
167
+ from k3_node.models import RWSEEncoder
168
+
169
+ rwse = np.random.rand(5, 16).astype("float32") # 16-step random-walk return probabilities
170
+ print(tuple(RWSEEncoder(num_rw_steps=16, pe_dim=8)(rwse).shape)) # (5, 8)
171
+ ```
172
+ """
173
+
174
+ def __init__(
175
+ self,
176
+ num_rw_steps: int = 16,
177
+ pe_dim: int = 20,
178
+ **kwargs,
179
+ ):
180
+ super().__init__(**kwargs)
181
+ self.num_rw_steps = num_rw_steps
182
+ self.pe_dim = pe_dim
183
+
184
+ self.raw_norm = layers.BatchNormalization(
185
+ axis=-1,
186
+ epsilon=1e-5,
187
+ momentum=0.9,
188
+ name="raw_norm",
189
+ )
190
+ self.pe_encoder = layers.Dense(pe_dim, use_bias=True, name="pe_encoder")
191
+
192
+ def build(self, input_shape=None):
193
+ if not self.raw_norm.built:
194
+ self.raw_norm.build((None, self.num_rw_steps))
195
+ if not self.pe_encoder.built:
196
+ self.pe_encoder.build((None, self.num_rw_steps))
197
+ super().build(input_shape)
198
+
199
+ def call(self, pestat_RWSE, training=False):
200
+ pe = self.raw_norm(pestat_RWSE, training=training)
201
+ return self.pe_encoder(pe)
202
+
203
+
204
+ class CustomGatedGCN(layers.Layer):
205
+ r"""Residual Gated Graph ConvNet layer with edge feature updates.
206
+
207
+ Reference:
208
+ "Residual Gated Graph ConvNets" (Bresson & Laurent, 2017).
209
+
210
+ Args:
211
+ in_dim (int): Input feature dimension.
212
+ out_dim (int): Output feature dimension.
213
+ dropout (float, optional): Dropout rate. (default: ``0.0``)
214
+ residual (bool, optional): Whether to use residual connections. (default: ``True``)
215
+ act (str, optional): Activation function. (default: ``"gelu"``)
216
+ **kwargs: Additional layer arguments.
217
+
218
+ Example:
219
+ ```python
220
+ import numpy as np
221
+ from k3_node.models import CustomGatedGCN
222
+
223
+ x = np.random.rand(5, 32).astype("float32") # node features
224
+ e = np.random.rand(6, 32).astype("float32") # edge features
225
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
226
+ batch = np.array([0, 0, 0, 1, 1]) # two graphs
227
+
228
+ layer = CustomGatedGCN(in_dim=32, out_dim=32, dropout=0.0, residual=True)
229
+ x_out, e_out = layer(x, edge_index, e) # updates node and edge features
230
+ print(tuple(x_out.shape), tuple(e_out.shape)) # (5, 32) (6, 32)
231
+ ```
232
+ """
233
+
234
+ def __init__(
235
+ self,
236
+ in_dim: int,
237
+ out_dim: int,
238
+ dropout: float = 0.0,
239
+ residual: bool = True,
240
+ act: str = "gelu",
241
+ **kwargs,
242
+ ):
243
+ super().__init__(**kwargs)
244
+ self.in_dim = in_dim
245
+ self.out_dim = out_dim
246
+ self.dropout_rate = dropout
247
+ self.residual = residual
248
+ self.act_name = act
249
+
250
+ self.A = layers.Dense(out_dim, use_bias=True, name="A")
251
+ self.B = layers.Dense(out_dim, use_bias=True, name="B")
252
+ self.C = layers.Dense(out_dim, use_bias=True, name="C")
253
+ self.D = layers.Dense(out_dim, use_bias=True, name="D")
254
+ self.E = layers.Dense(out_dim, use_bias=True, name="E")
255
+
256
+ self.bn_node_x = layers.BatchNormalization(
257
+ axis=-1, epsilon=1e-5, momentum=0.9, name="bn_node_x"
258
+ )
259
+ self.bn_edge_e = layers.BatchNormalization(
260
+ axis=-1, epsilon=1e-5, momentum=0.9, name="bn_edge_e"
261
+ )
262
+ self.dropout_node = layers.Dropout(dropout)
263
+ self.dropout_edge = layers.Dropout(dropout)
264
+
265
+ def build(self, input_shape=None):
266
+ if not self.A.built:
267
+ self.A.build((None, self.in_dim))
268
+ self.B.build((None, self.in_dim))
269
+ self.C.build((None, self.in_dim))
270
+ self.D.build((None, self.in_dim))
271
+ self.E.build((None, self.in_dim))
272
+ self.bn_node_x.build((None, self.out_dim))
273
+ self.bn_edge_e.build((None, self.out_dim))
274
+ super().build(input_shape)
275
+
276
+ def call(self, x, edge_index, edge_attr, training=False):
277
+ x_in = x
278
+ e_in = edge_attr
279
+
280
+ Ax = self.A(x)
281
+ Bx = self.B(x)
282
+ Ce = self.C(edge_attr)
283
+ Dx = self.D(x)
284
+ Ex = self.E(x)
285
+
286
+ src = ops.cast(edge_index[0], "int32")
287
+ dst = ops.cast(edge_index[1], "int32")
288
+
289
+ Dx_i = ops.take(Dx, dst, axis=0)
290
+ Ex_j = ops.take(Ex, src, axis=0)
291
+ Bx_j = ops.take(Bx, src, axis=0)
292
+
293
+ e_ij = Dx_i + Ex_j + Ce
294
+ sigma_ij = ops.sigmoid(e_ij)
295
+
296
+ num_nodes = ops.shape(x)[0]
297
+ sum_sigma_x = segment_sum(sigma_ij * Bx_j, dst, num_segments=num_nodes)
298
+ sum_sigma = segment_sum(sigma_ij, dst, num_segments=num_nodes)
299
+ aggr_out = sum_sigma_x / (sum_sigma + 1e-6)
300
+
301
+ x_out = self.bn_node_x(Ax + aggr_out, training=training)
302
+ e_out = self.bn_edge_e(e_ij, training=training)
303
+
304
+ act_fn = _get_act(self.act_name)
305
+ x_out = act_fn(x_out)
306
+ e_out = act_fn(e_out)
307
+
308
+ x_out = self.dropout_node(x_out, training=training)
309
+ e_out = self.dropout_edge(e_out, training=training)
310
+
311
+ if self.residual:
312
+ x_out = x_in + x_out
313
+ e_out = e_in + e_out
314
+
315
+ return x_out, e_out
316
+
317
+
318
+ class TransformerSelfAttention(layers.Layer):
319
+ r"""Multi-Head Self-Attention layer matching PyTorch `nn.MultiheadAttention`.
320
+
321
+ Args:
322
+ embed_dim (int): Total embedding dimension.
323
+ num_heads (int): Number of attention heads.
324
+ dropout (float, optional): Attention dropout rate. (default: ``0.0``)
325
+ **kwargs: Additional layer arguments.
326
+ """
327
+
328
+ def __init__(
329
+ self,
330
+ embed_dim: int,
331
+ num_heads: int,
332
+ dropout: float = 0.0,
333
+ **kwargs,
334
+ ):
335
+ super().__init__(**kwargs)
336
+ self.supports_masking = True
337
+ self.embed_dim = embed_dim
338
+ self.num_heads = num_heads
339
+ self.head_dim = embed_dim // num_heads
340
+ self.scaling = 1.0 / math.sqrt(self.head_dim)
341
+ self.dropout_rate = dropout
342
+
343
+ self.q_proj = layers.Dense(embed_dim, use_bias=True, name="q_proj")
344
+ self.k_proj = layers.Dense(embed_dim, use_bias=True, name="k_proj")
345
+ self.v_proj = layers.Dense(embed_dim, use_bias=True, name="v_proj")
346
+ self.out_proj = layers.Dense(embed_dim, use_bias=True, name="out_proj")
347
+ self.dropout = layers.Dropout(dropout)
348
+
349
+ def build(self, input_shape=None):
350
+ if not self.q_proj.built:
351
+ self.q_proj.build((None, self.embed_dim))
352
+ self.k_proj.build((None, self.embed_dim))
353
+ self.v_proj.build((None, self.embed_dim))
354
+ self.out_proj.build((None, self.embed_dim))
355
+ super().build(input_shape)
356
+
357
+ def call(self, x, mask=None, training=False):
358
+ shape = ops.shape(x)
359
+ bsz, n_node = shape[0], shape[1]
360
+
361
+ q = self.q_proj(x)
362
+ k = self.k_proj(x)
363
+ v = self.v_proj(x)
364
+
365
+ q = ops.transpose(
366
+ ops.reshape(q, (bsz, n_node, self.num_heads, self.head_dim)),
367
+ (0, 2, 1, 3),
368
+ ) * self.scaling
369
+ k = ops.transpose(
370
+ ops.reshape(k, (bsz, n_node, self.num_heads, self.head_dim)),
371
+ (0, 2, 1, 3),
372
+ )
373
+ v = ops.transpose(
374
+ ops.reshape(v, (bsz, n_node, self.num_heads, self.head_dim)),
375
+ (0, 2, 1, 3),
376
+ )
377
+
378
+ scores = ops.matmul(q, ops.transpose(k, (0, 1, 3, 2))) # [B, H, N, N]
379
+
380
+ if mask is not None:
381
+ # mask: [B, N] boolean tensor, True for valid tokens, False for padding
382
+ # PyTorch key_padding_mask: True where padding
383
+ key_padding_mask = ~mask
384
+ scores = ops.where(
385
+ ops.expand_dims(ops.expand_dims(key_padding_mask, axis=1), axis=2),
386
+ float("-inf"),
387
+ scores,
388
+ )
389
+
390
+ attn_weights = ops.softmax(scores, axis=-1)
391
+ # Replace NaNs from all-inf rows
392
+ attn_weights = ops.where(ops.isnan(attn_weights), 0.0, attn_weights)
393
+ attn_weights = self.dropout(attn_weights, training=training)
394
+
395
+ attn = ops.matmul(attn_weights, v) # [B, H, N, head_dim]
396
+ attn = ops.transpose(attn, (0, 2, 1, 3))
397
+ attn = ops.reshape(attn, (bsz, n_node, self.embed_dim))
398
+
399
+ return self.out_proj(attn)
400
+
401
+
402
+ class GPSLayer(layers.Layer):
403
+ r"""GraphGPS hybrid layer combining local MPNN (e.g. `CustomGatedGCN`) and global Multi-Head Attention.
404
+
405
+ Reference:
406
+ "Recipe for a General, Powerful, Scalable Graph Transformer" (NeurIPS 2022).
407
+
408
+ Args:
409
+ dim_h (int): Hidden embedding dimension.
410
+ local_gnn_type (str, optional): Local MPNN type. (default: ``"CustomGatedGCN"``)
411
+ global_model_type (str, optional): Global attention type. (default: ``"Transformer"``)
412
+ num_heads (int, optional): Number of attention heads. (default: ``8``)
413
+ act (str, optional): Activation function. (default: ``"gelu"``)
414
+ dropout (float, optional): Dropout rate. (default: ``0.0``)
415
+ attn_dropout (float, optional): Attention dropout rate. (default: ``0.0``)
416
+ layer_norm (bool, optional): Whether to use LayerNorm. (default: ``False``)
417
+ batch_norm (bool, optional): Whether to use BatchNorm. (default: ``True``)
418
+ **kwargs: Additional layer arguments.
419
+
420
+ Example:
421
+ ```python
422
+ import numpy as np
423
+ from k3_node.models import GPSLayer
424
+
425
+ x = np.random.rand(5, 32).astype("float32") # node features
426
+ e = np.random.rand(6, 32).astype("float32") # edge features
427
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
428
+ batch = np.array([0, 0, 0, 1, 1]) # two graphs
429
+
430
+ # Local message passing plus global attention over each graph's nodes
431
+ layer = GPSLayer(dim_h=32, local_gnn_type="CustomGatedGCN", global_model_type="Transformer", num_heads=4)
432
+ x_out, e_out = layer(x, edge_index, e, batch=batch)
433
+ print(tuple(x_out.shape), tuple(e_out.shape)) # (5, 32) (6, 32)
434
+ ```
435
+ """
436
+
437
+ def __init__(
438
+ self,
439
+ dim_h: int,
440
+ local_gnn_type: str = "CustomGatedGCN",
441
+ global_model_type: str = "Transformer",
442
+ num_heads: int = 8,
443
+ act: str = "gelu",
444
+ dropout: float = 0.0,
445
+ attn_dropout: float = 0.0,
446
+ layer_norm: bool = False,
447
+ batch_norm: bool = True,
448
+ **kwargs,
449
+ ):
450
+ super().__init__(**kwargs)
451
+ self.dim_h = dim_h
452
+ self.local_gnn_type = local_gnn_type
453
+ self.global_model_type = global_model_type
454
+ self.num_heads = num_heads
455
+ self.act_name = act
456
+ self.dropout_rate = dropout
457
+ self.attn_dropout_rate = attn_dropout
458
+ self.layer_norm = layer_norm
459
+ self.batch_norm = batch_norm
460
+
461
+ if local_gnn_type == "CustomGatedGCN":
462
+ self.local_model = CustomGatedGCN(
463
+ in_dim=dim_h,
464
+ out_dim=dim_h,
465
+ dropout=dropout,
466
+ residual=True,
467
+ act=act,
468
+ name="local_model",
469
+ )
470
+ else:
471
+ self.local_model = None
472
+
473
+ if global_model_type in ["Transformer", "BiasedTransformer"]:
474
+ self.self_attn = TransformerSelfAttention(
475
+ embed_dim=dim_h,
476
+ num_heads=num_heads,
477
+ dropout=attn_dropout,
478
+ name="self_attn",
479
+ )
480
+ else:
481
+ self.self_attn = None
482
+
483
+ if layer_norm:
484
+ self.norm1_local = layers.LayerNormalization(epsilon=1e-5, name="norm1_local")
485
+ self.norm1_attn = layers.LayerNormalization(epsilon=1e-5, name="norm1_attn")
486
+ self.norm2 = layers.LayerNormalization(epsilon=1e-5, name="norm2")
487
+ elif batch_norm:
488
+ self.norm1_local = layers.BatchNormalization(
489
+ axis=-1, epsilon=1e-5, momentum=0.9, name="norm1_local"
490
+ )
491
+ self.norm1_attn = layers.BatchNormalization(
492
+ axis=-1, epsilon=1e-5, momentum=0.9, name="norm1_attn"
493
+ )
494
+ self.norm2 = layers.BatchNormalization(
495
+ axis=-1, epsilon=1e-5, momentum=0.9, name="norm2"
496
+ )
497
+ else:
498
+ self.norm1_local = None
499
+ self.norm1_attn = None
500
+ self.norm2 = None
501
+
502
+ self.dropout_local = layers.Dropout(dropout)
503
+ self.dropout_attn = layers.Dropout(dropout)
504
+
505
+ self.ff_linear1 = layers.Dense(dim_h * 2, use_bias=True, name="ff_linear1")
506
+ self.ff_linear2 = layers.Dense(dim_h, use_bias=True, name="ff_linear2")
507
+ self.ff_dropout1 = layers.Dropout(dropout)
508
+ self.ff_dropout2 = layers.Dropout(dropout)
509
+
510
+ def build(self, input_shape=None):
511
+ if self.local_model is not None and not self.local_model.built:
512
+ self.local_model.build(None)
513
+ if self.self_attn is not None and not self.self_attn.built:
514
+ self.self_attn.build(None)
515
+ if self.norm1_local is not None and not self.norm1_local.built:
516
+ self.norm1_local.build((None, self.dim_h))
517
+ if self.norm1_attn is not None and not self.norm1_attn.built:
518
+ self.norm1_attn.build((None, self.dim_h))
519
+ if not self.ff_linear1.built:
520
+ self.ff_linear1.build((None, self.dim_h))
521
+ self.ff_linear2.build((None, self.dim_h * 2))
522
+ if self.norm2 is not None and not self.norm2.built:
523
+ self.norm2.build((None, self.dim_h))
524
+ super().build(input_shape)
525
+
526
+ def call(self, x, edge_index, edge_attr, batch=None, training=False):
527
+ h_in1 = x
528
+ h_out_list = []
529
+
530
+ # Local MPNN
531
+ if self.local_model is not None:
532
+ h_local, edge_attr = self.local_model(
533
+ x, edge_index, edge_attr, training=training
534
+ )
535
+ # CustomGatedGCN handles residual internally
536
+ if self.norm1_local is not None:
537
+ h_local = self.norm1_local(h_local, training=training)
538
+ h_out_list.append(h_local)
539
+
540
+ # Global Attention
541
+ if self.self_attn is not None:
542
+ if batch is None:
543
+ h_dense = ops.expand_dims(x, axis=0)
544
+ mask = ops.ones((1, ops.shape(x)[0]), dtype="bool")
545
+ h_attn = self.self_attn(h_dense, mask=mask, training=training)[0]
546
+ else:
547
+ batch_np = ops.convert_to_numpy(batch)
548
+ B_int = int(batch_np.max()) + 1 if len(batch_np) > 0 else 1
549
+ counts = np.bincount(batch_np, minlength=B_int)
550
+ max_nodes = int(counts.max()) if len(counts) > 0 else 0
551
+
552
+ offsets = np.zeros(len(batch_np), dtype=np.int32)
553
+ curr = np.zeros(B_int, dtype=np.int32)
554
+ for i, b in enumerate(batch_np):
555
+ offsets[i] = curr[b]
556
+ curr[b] += 1
557
+
558
+ offsets_t = ops.convert_to_tensor(offsets, dtype="int32")
559
+ batch_cast = ops.cast(batch, "int32")
560
+ indices = ops.stack([batch_cast, offsets_t], axis=1)
561
+
562
+ dense_x = ops.scatter_update(
563
+ ops.zeros((B_int, max_nodes, self.dim_h), dtype=x.dtype),
564
+ indices,
565
+ x,
566
+ )
567
+
568
+ mask_np = np.zeros((B_int, max_nodes), dtype=bool)
569
+ for b, count in enumerate(counts):
570
+ mask_np[b, :count] = True
571
+ mask = ops.convert_to_tensor(mask_np)
572
+
573
+ h_attn_dense = self.self_attn(dense_x, mask=mask, training=training)
574
+
575
+ flat_attn = ops.reshape(h_attn_dense, (B_int * max_nodes, self.dim_h))
576
+ flat_indices = batch_cast * max_nodes + offsets_t
577
+ h_attn = ops.take(flat_attn, flat_indices, axis=0)
578
+
579
+ h_attn = self.dropout_attn(h_attn, training=training)
580
+ h_attn = h_in1 + h_attn
581
+ if self.norm1_attn is not None:
582
+ h_attn = self.norm1_attn(h_attn, training=training)
583
+ h_out_list.append(h_attn)
584
+
585
+ # Sum local and global representations
586
+ h = sum(h_out_list)
587
+
588
+ # Feed Forward block
589
+ act_fn = _get_act(self.act_name)
590
+ ff_out = self.ff_dropout1(act_fn(self.ff_linear1(h)), training=training)
591
+ ff_out = self.ff_dropout2(self.ff_linear2(ff_out), training=training)
592
+ h = h + ff_out
593
+
594
+ if self.norm2 is not None:
595
+ h = self.norm2(h, training=training)
596
+
597
+ return h, edge_attr
598
+
599
+
600
+ class SANGraphHead(layers.Layer):
601
+ r"""Prediction head for graph-level tasks from the Spectral Attention Network (SAN).
602
+
603
+ Args:
604
+ dim_in (int): Input feature dimension.
605
+ dim_out (int): Output feature dimension. (default: ``1``)
606
+ L (int, optional): Number of hidden layers. (default: ``2``)
607
+ act (str, optional): Activation function. (default: ``"gelu"``)
608
+ pooling (str, optional): Graph pooling method ('mean', 'add', 'max'). (default: ``"mean"``)
609
+ **kwargs: Additional layer arguments.
610
+
611
+ Example:
612
+ ```python
613
+ import numpy as np
614
+ from k3_node.models import SANGraphHead
615
+
616
+ x = np.random.rand(5, 32).astype("float32") # node features
617
+ e = np.random.rand(6, 32).astype("float32") # edge features
618
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
619
+ batch = np.array([0, 0, 0, 1, 1]) # two graphs
620
+
621
+ head = SANGraphHead(dim_in=32, dim_out=1, L=2, pooling="mean")
622
+ print(tuple(head(x, batch=batch).shape)) # (2, 1): one prediction per graph
623
+ ```
624
+ """
625
+
626
+ def __init__(
627
+ self,
628
+ dim_in: int,
629
+ dim_out: int = 1,
630
+ L: int = 2,
631
+ act: str = "gelu",
632
+ pooling: str = "mean",
633
+ **kwargs,
634
+ ):
635
+ super().__init__(**kwargs)
636
+ self.dim_in = dim_in
637
+ self.dim_out = dim_out
638
+ self.L = L
639
+ self.act_name = act
640
+ self.pooling = pooling.lower()
641
+
642
+ self.FC_layers = []
643
+ for l in range(L):
644
+ in_dim_l = dim_in // (2**l)
645
+ out_dim_l = dim_in // (2 ** (l + 1))
646
+ self.FC_layers.append(
647
+ layers.Dense(out_dim_l, use_bias=True, name=f"FC_layers_{l}")
648
+ )
649
+ self.FC_layers.append(
650
+ layers.Dense(dim_out, use_bias=True, name=f"FC_layers_{L}")
651
+ )
652
+
653
+ def build(self, input_shape=None):
654
+ curr_dim = self.dim_in
655
+ for layer in self.FC_layers:
656
+ if not layer.built:
657
+ layer.build((None, curr_dim))
658
+ curr_dim = layer.units
659
+ super().build(input_shape)
660
+
661
+ def call(self, x, batch=None, training=False, batch_size=None):
662
+ if self.pooling in ["mean", "avg"]:
663
+ graph_emb = global_mean_pool(x, batch, size=batch_size)
664
+ elif self.pooling in ["add", "sum"]:
665
+ graph_emb = global_add_pool(x, batch, size=batch_size)
666
+ elif self.pooling == "max":
667
+ graph_emb = global_max_pool(x, batch, size=batch_size)
668
+ else:
669
+ raise ValueError(f"Unknown pooling method '{self.pooling}'")
670
+
671
+ act_fn = _get_act(self.act_name)
672
+ for l in range(self.L):
673
+ graph_emb = self.FC_layers[l](graph_emb)
674
+ graph_emb = act_fn(graph_emb)
675
+
676
+ graph_emb = self.FC_layers[self.L](graph_emb)
677
+ return graph_emb
678
+
679
+
680
+ class GPSModel(keras.Model):
681
+ r"""GraphGPS: General Powerful Scalable Graph Transformer from the
682
+ `"Recipe for a General, Powerful, Scalable Graph Transformer"
683
+ <https://arxiv.org/abs/2205.12454>`_ paper (NeurIPS 2022).
684
+
685
+ Args:
686
+ dim_in (int, optional): Initial input feature dimension. (default: ``256``)
687
+ dim_out (int, optional): Target output dimension. (default: ``1``)
688
+ num_layers (int, optional): Number of GPS layers. (default: ``16``)
689
+ dim_hidden (int, optional): Hidden embedding dimension. (default: ``256``)
690
+ num_heads (int, optional): Number of attention heads. (default: ``8``)
691
+ local_gnn_type (str, optional): Local MPNN layer type. (default: ``"CustomGatedGCN"``)
692
+ act (str, optional): Activation function. (default: ``"gelu"``)
693
+ dropout (float, optional): Dropout probability. (default: ``0.1``)
694
+ attn_dropout (float, optional): Attention dropout probability. (default: ``0.1``)
695
+ batch_norm (bool, optional): Whether to use batch normalization. (default: ``True``)
696
+ layer_norm (bool, optional): Whether to use layer normalization. (default: ``False``)
697
+ node_encoder_type (str, optional): Node encoder type ("Atom+RWSE", "Atom", "Linear", or None). (default: ``"Atom+RWSE"``)
698
+ edge_encoder_type (str, optional): Edge encoder type ("Bond", "Linear", or None). (default: ``"Bond"``)
699
+ atom_feature_dims (List[int], optional): Categorical feature vocabulary sizes for atom features.
700
+ bond_feature_dims (List[int], optional): Categorical feature vocabulary sizes for bond features.
701
+ rwse_num_steps (int, optional): Number of RWSE steps. (default: ``16``)
702
+ rwse_dim_pe (int, optional): RWSE embedding dimension. (default: ``20``)
703
+ graph_pooling (str, optional): Graph pooling type ('mean', 'add', 'max'). (default: ``"mean"``)
704
+ head_layers (int, optional): Number of hidden layers in prediction head. (default: ``2``)
705
+ **kwargs: Additional model arguments.
706
+
707
+ Example:
708
+ ```python
709
+ import numpy as np
710
+ from k3_node.models import GPSModel
711
+
712
+ # OGB-style integer atom and bond features (9 per atom, 3 per bond)
713
+ x = np.random.randint(0, 2, size=(5, 9))
714
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
715
+ edge_attr = np.random.randint(0, 2, size=(6, 3))
716
+ batch = np.array([0, 0, 0, 1, 1]) # two graphs
717
+ rwse = np.random.rand(5, 16).astype("float32") # random-walk structural encodings
718
+
719
+ model = GPSModel(dim_in=32, dim_out=1, num_layers=2, dim_hidden=32, num_heads=4,
720
+ node_encoder_type="Atom+RWSE", edge_encoder_type="Bond",
721
+ rwse_num_steps=16, rwse_dim_pe=8)
722
+ pred = model(x, edge_index, edge_attr=edge_attr, pestat_RWSE=rwse, batch=batch, batch_size=2)
723
+ print(tuple(pred.shape)) # (2, 1): one prediction per graph
724
+ ```
725
+ """
726
+
727
+ def __init__(
728
+ self,
729
+ dim_in: int = 256,
730
+ dim_out: int = 1,
731
+ num_layers: int = 16,
732
+ dim_hidden: int = 256,
733
+ num_heads: int = 8,
734
+ local_gnn_type: str = "CustomGatedGCN",
735
+ act: str = "gelu",
736
+ dropout: float = 0.1,
737
+ attn_dropout: float = 0.1,
738
+ batch_norm: bool = True,
739
+ layer_norm: bool = False,
740
+ node_encoder_type: Optional[str] = "Atom+RWSE",
741
+ edge_encoder_type: Optional[str] = "Bond",
742
+ atom_feature_dims: Optional[List[int]] = None,
743
+ bond_feature_dims: Optional[List[int]] = None,
744
+ rwse_num_steps: int = 16,
745
+ rwse_dim_pe: int = 20,
746
+ graph_pooling: str = "mean",
747
+ head_layers: int = 2,
748
+ **kwargs,
749
+ ):
750
+ super().__init__(**kwargs)
751
+ self.dim_in = dim_in
752
+ self.dim_out = dim_out
753
+ self.num_layers = num_layers
754
+ self.dim_hidden = dim_hidden
755
+ self.num_heads = num_heads
756
+ self.local_gnn_type = local_gnn_type
757
+ self.act_name = act
758
+ self.dropout_rate = dropout
759
+ self.attn_dropout_rate = attn_dropout
760
+ self.batch_norm = batch_norm
761
+ self.layer_norm = layer_norm
762
+ self.node_encoder_type = node_encoder_type
763
+ self.edge_encoder_type = edge_encoder_type
764
+ self.rwse_num_steps = rwse_num_steps
765
+ self.rwse_dim_pe = rwse_dim_pe
766
+ self.graph_pooling = graph_pooling
767
+ self.head_layers = head_layers
768
+
769
+ # Node Encoder
770
+ if node_encoder_type == "Atom+RWSE":
771
+ self.atom_encoder = AtomEncoder(
772
+ emb_dim=dim_hidden - rwse_dim_pe,
773
+ feature_dims=atom_feature_dims,
774
+ name="atom_encoder",
775
+ )
776
+ self.rwse_encoder = RWSEEncoder(
777
+ num_rw_steps=rwse_num_steps,
778
+ pe_dim=rwse_dim_pe,
779
+ name="rwse_encoder",
780
+ )
781
+ elif node_encoder_type == "Atom":
782
+ self.atom_encoder = AtomEncoder(
783
+ emb_dim=dim_hidden,
784
+ feature_dims=atom_feature_dims,
785
+ name="atom_encoder",
786
+ )
787
+ self.rwse_encoder = None
788
+ elif node_encoder_type == "Linear":
789
+ self.linear_node_encoder = layers.Dense(dim_hidden, name="linear_node_encoder")
790
+ self.atom_encoder = None
791
+ self.rwse_encoder = None
792
+ else:
793
+ self.atom_encoder = None
794
+ self.rwse_encoder = None
795
+ self.linear_node_encoder = None
796
+
797
+ # Edge Encoder
798
+ if edge_encoder_type == "Bond":
799
+ self.bond_encoder = BondEncoder(
800
+ emb_dim=dim_hidden,
801
+ feature_dims=bond_feature_dims,
802
+ name="bond_encoder",
803
+ )
804
+ elif edge_encoder_type == "Linear":
805
+ self.linear_edge_encoder = layers.Dense(dim_hidden, name="linear_edge_encoder")
806
+ self.bond_encoder = None
807
+ else:
808
+ self.bond_encoder = None
809
+ self.linear_edge_encoder = None
810
+
811
+ # GPS Layers
812
+ self.gps_layers = [
813
+ GPSLayer(
814
+ dim_h=dim_hidden,
815
+ local_gnn_type=local_gnn_type,
816
+ global_model_type="Transformer",
817
+ num_heads=num_heads,
818
+ act=act,
819
+ dropout=dropout,
820
+ attn_dropout=attn_dropout,
821
+ layer_norm=layer_norm,
822
+ batch_norm=batch_norm,
823
+ name=f"gps_layer_{i}",
824
+ )
825
+ for i in range(num_layers)
826
+ ]
827
+
828
+ # SANGraphHead
829
+ self.post_mp = SANGraphHead(
830
+ dim_in=dim_hidden,
831
+ dim_out=dim_out,
832
+ L=head_layers,
833
+ act=act,
834
+ pooling=graph_pooling,
835
+ name="post_mp",
836
+ )
837
+
838
+ def build(self, input_shape=None):
839
+ if self.atom_encoder is not None and not self.atom_encoder.built:
840
+ self.atom_encoder.build(None)
841
+ if self.rwse_encoder is not None and not self.rwse_encoder.built:
842
+ self.rwse_encoder.build(None)
843
+ if self.bond_encoder is not None and not self.bond_encoder.built:
844
+ self.bond_encoder.build(None)
845
+ for layer in self.gps_layers:
846
+ if not layer.built:
847
+ layer.build(None)
848
+ if not self.post_mp.built:
849
+ self.post_mp.build(None)
850
+ super().build(input_shape)
851
+
852
+ def call(
853
+ self,
854
+ x,
855
+ edge_index,
856
+ edge_attr=None,
857
+ pestat_RWSE=None,
858
+ batch=None,
859
+ training=False,
860
+ batch_size=None,
861
+ ):
862
+ # Node encoding
863
+ if self.node_encoder_type == "Atom+RWSE":
864
+ x_emb = self.atom_encoder(x)
865
+ if pestat_RWSE is not None and self.rwse_encoder is not None:
866
+ pe_emb = self.rwse_encoder(pestat_RWSE, training=training)
867
+ x = ops.concatenate([x_emb, pe_emb], axis=-1)
868
+ else:
869
+ x = x_emb
870
+ elif self.node_encoder_type == "Atom":
871
+ x = self.atom_encoder(x)
872
+ elif self.node_encoder_type == "Linear" and self.linear_node_encoder is not None:
873
+ x = self.linear_node_encoder(x)
874
+
875
+ # Edge encoding
876
+ if self.edge_encoder_type == "Bond" and edge_attr is not None:
877
+ edge_attr = self.bond_encoder(edge_attr)
878
+ elif self.edge_encoder_type == "Linear" and self.linear_edge_encoder is not None:
879
+ edge_attr = self.linear_edge_encoder(edge_attr)
880
+
881
+ # GPS layers
882
+ for layer in self.gps_layers:
883
+ x, edge_attr = layer(
884
+ x, edge_index, edge_attr, batch=batch, training=training
885
+ )
886
+
887
+ # Head
888
+ pred = self.post_mp(x, batch=batch, training=training, batch_size=batch_size)
889
+ return pred
890
+
891
+
892
+ def load_gps_weights(model: GPSModel, checkpoint_path: str):
893
+ r"""Loads trained PyTorch GraphGPS checkpoint weights into a Keras 3 `GPSModel`.
894
+
895
+ Args:
896
+ model (GPSModel): Target `GPSModel` instance.
897
+ checkpoint_path (str): Path to PyTorch `.ckpt` or `.pt` checkpoint file.
898
+
899
+ Returns:
900
+ GPSModel: The model with loaded weights.
901
+ """
902
+ import torch
903
+
904
+ ckpt = torch.load(checkpoint_path, map_location="cpu")
905
+ if "model_state" in ckpt:
906
+ state_dict = ckpt["model_state"]
907
+ elif "state_dict" in ckpt:
908
+ state_dict = ckpt["state_dict"]
909
+ else:
910
+ state_dict = ckpt
911
+
912
+ clean_dict = {}
913
+ for k, v in state_dict.items():
914
+ if k.startswith("model."):
915
+ k = k[6:]
916
+ clean_dict[k] = v
917
+
918
+ def _to_tensor(t):
919
+ arr = t.detach().cpu().numpy()
920
+ return ops.convert_to_tensor(arr, dtype="float32")
921
+
922
+ # Build model if needed
923
+ if not model.built:
924
+ model.build(None)
925
+
926
+ # 1. Node Encoder
927
+ if model.node_encoder_type in ["Atom+RWSE", "Atom"] and model.atom_encoder is not None:
928
+ for i, emb in enumerate(model.atom_encoder.atom_embedding_list):
929
+ k = f"encoder.node_encoder.encoder1.atom_embedding_list.{i}.weight"
930
+ if k in clean_dict:
931
+ emb.embeddings.assign(_to_tensor(clean_dict[k]))
932
+
933
+ if model.node_encoder_type == "Atom+RWSE" and model.rwse_encoder is not None:
934
+ p = "encoder.node_encoder.encoder2"
935
+ if f"{p}.raw_norm.weight" in clean_dict:
936
+ model.rwse_encoder.raw_norm.gamma.assign(_to_tensor(clean_dict[f"{p}.raw_norm.weight"]))
937
+ if f"{p}.raw_norm.bias" in clean_dict:
938
+ model.rwse_encoder.raw_norm.beta.assign(_to_tensor(clean_dict[f"{p}.raw_norm.bias"]))
939
+ if f"{p}.raw_norm.running_mean" in clean_dict:
940
+ model.rwse_encoder.raw_norm.moving_mean.assign(_to_tensor(clean_dict[f"{p}.raw_norm.running_mean"]))
941
+ if f"{p}.raw_norm.running_var" in clean_dict:
942
+ model.rwse_encoder.raw_norm.moving_variance.assign(_to_tensor(clean_dict[f"{p}.raw_norm.running_var"]))
943
+
944
+ if f"{p}.pe_encoder.weight" in clean_dict:
945
+ model.rwse_encoder.pe_encoder.kernel.assign(_to_tensor(clean_dict[f"{p}.pe_encoder.weight"].t()))
946
+ if f"{p}.pe_encoder.bias" in clean_dict:
947
+ model.rwse_encoder.pe_encoder.bias.assign(_to_tensor(clean_dict[f"{p}.pe_encoder.bias"]))
948
+
949
+ # 2. Edge Encoder
950
+ if model.edge_encoder_type == "Bond" and model.bond_encoder is not None:
951
+ for i, emb in enumerate(model.bond_encoder.bond_embedding_list):
952
+ k = f"encoder.edge_encoder.bond_embedding_list.{i}.weight"
953
+ if k in clean_dict:
954
+ emb.embeddings.assign(_to_tensor(clean_dict[k]))
955
+
956
+ # 3. GPS Layers
957
+ dim_h = model.dim_hidden
958
+ for i, gps_layer in enumerate(model.gps_layers):
959
+ p = f"layers.{i}"
960
+
961
+ # Local MPNN: CustomGatedGCN
962
+ if gps_layer.local_model is not None:
963
+ lm = gps_layer.local_model
964
+ for proj_name in ["A", "B", "C", "D", "E"]:
965
+ dense_proj = getattr(lm, proj_name)
966
+ if f"{p}.local_model.{proj_name}.weight" in clean_dict:
967
+ dense_proj.kernel.assign(_to_tensor(clean_dict[f"{p}.local_model.{proj_name}.weight"].t()))
968
+ if f"{p}.local_model.{proj_name}.bias" in clean_dict:
969
+ dense_proj.bias.assign(_to_tensor(clean_dict[f"{p}.local_model.{proj_name}.bias"]))
970
+
971
+ for bn_name, bn_layer in [("bn_node_x", lm.bn_node_x), ("bn_edge_e", lm.bn_edge_e)]:
972
+ if f"{p}.local_model.{bn_name}.weight" in clean_dict:
973
+ bn_layer.gamma.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.weight"]))
974
+ if f"{p}.local_model.{bn_name}.bias" in clean_dict:
975
+ bn_layer.beta.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.bias"]))
976
+ if f"{p}.local_model.{bn_name}.running_mean" in clean_dict:
977
+ bn_layer.moving_mean.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.running_mean"]))
978
+ if f"{p}.local_model.{bn_name}.running_var" in clean_dict:
979
+ bn_layer.moving_variance.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.running_var"]))
980
+
981
+ if gps_layer.norm1_local is not None:
982
+ if f"{p}.norm1_local.weight" in clean_dict:
983
+ gps_layer.norm1_local.gamma.assign(_to_tensor(clean_dict[f"{p}.norm1_local.weight"]))
984
+ if f"{p}.norm1_local.bias" in clean_dict:
985
+ gps_layer.norm1_local.beta.assign(_to_tensor(clean_dict[f"{p}.norm1_local.bias"]))
986
+ if hasattr(gps_layer.norm1_local, "moving_mean") and f"{p}.norm1_local.running_mean" in clean_dict:
987
+ gps_layer.norm1_local.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm1_local.running_mean"]))
988
+ if hasattr(gps_layer.norm1_local, "moving_variance") and f"{p}.norm1_local.running_var" in clean_dict:
989
+ gps_layer.norm1_local.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm1_local.running_var"]))
990
+
991
+ # Global Attention: MultiHeadAttention
992
+ if gps_layer.self_attn is not None:
993
+ sa = gps_layer.self_attn
994
+ if f"{p}.self_attn.in_proj_weight" in clean_dict:
995
+ in_w = clean_dict[f"{p}.self_attn.in_proj_weight"]
996
+ sa.q_proj.kernel.assign(_to_tensor(in_w[:dim_h, :].t()))
997
+ sa.k_proj.kernel.assign(_to_tensor(in_w[dim_h:2*dim_h, :].t()))
998
+ sa.v_proj.kernel.assign(_to_tensor(in_w[2*dim_h:, :].t()))
999
+
1000
+ if f"{p}.self_attn.in_proj_bias" in clean_dict:
1001
+ in_b = clean_dict[f"{p}.self_attn.in_proj_bias"]
1002
+ sa.q_proj.bias.assign(_to_tensor(in_b[:dim_h]))
1003
+ sa.k_proj.bias.assign(_to_tensor(in_b[dim_h:2*dim_h]))
1004
+ sa.v_proj.bias.assign(_to_tensor(in_b[2*dim_h:]))
1005
+
1006
+ if f"{p}.self_attn.out_proj.weight" in clean_dict:
1007
+ sa.out_proj.kernel.assign(_to_tensor(clean_dict[f"{p}.self_attn.out_proj.weight"].t()))
1008
+ if f"{p}.self_attn.out_proj.bias" in clean_dict:
1009
+ sa.out_proj.bias.assign(_to_tensor(clean_dict[f"{p}.self_attn.out_proj.bias"]))
1010
+
1011
+ if gps_layer.norm1_attn is not None:
1012
+ if f"{p}.norm1_attn.weight" in clean_dict:
1013
+ gps_layer.norm1_attn.gamma.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.weight"]))
1014
+ if f"{p}.norm1_attn.bias" in clean_dict:
1015
+ gps_layer.norm1_attn.beta.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.bias"]))
1016
+ if hasattr(gps_layer.norm1_attn, "moving_mean") and f"{p}.norm1_attn.running_mean" in clean_dict:
1017
+ gps_layer.norm1_attn.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.running_mean"]))
1018
+ if hasattr(gps_layer.norm1_attn, "moving_variance") and f"{p}.norm1_attn.running_var" in clean_dict:
1019
+ gps_layer.norm1_attn.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.running_var"]))
1020
+
1021
+ # FFN
1022
+ if f"{p}.ff_linear1.weight" in clean_dict:
1023
+ gps_layer.ff_linear1.kernel.assign(_to_tensor(clean_dict[f"{p}.ff_linear1.weight"].t()))
1024
+ if f"{p}.ff_linear1.bias" in clean_dict:
1025
+ gps_layer.ff_linear1.bias.assign(_to_tensor(clean_dict[f"{p}.ff_linear1.bias"]))
1026
+ if f"{p}.ff_linear2.weight" in clean_dict:
1027
+ gps_layer.ff_linear2.kernel.assign(_to_tensor(clean_dict[f"{p}.ff_linear2.weight"].t()))
1028
+ if f"{p}.ff_linear2.bias" in clean_dict:
1029
+ gps_layer.ff_linear2.bias.assign(_to_tensor(clean_dict[f"{p}.ff_linear2.bias"]))
1030
+
1031
+ if gps_layer.norm2 is not None:
1032
+ if f"{p}.norm2.weight" in clean_dict:
1033
+ gps_layer.norm2.gamma.assign(_to_tensor(clean_dict[f"{p}.norm2.weight"]))
1034
+ if f"{p}.norm2.bias" in clean_dict:
1035
+ gps_layer.norm2.beta.assign(_to_tensor(clean_dict[f"{p}.norm2.bias"]))
1036
+ if hasattr(gps_layer.norm2, "moving_mean") and f"{p}.norm2.running_mean" in clean_dict:
1037
+ gps_layer.norm2.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm2.running_mean"]))
1038
+ if hasattr(gps_layer.norm2, "moving_variance") and f"{p}.norm2.running_var" in clean_dict:
1039
+ gps_layer.norm2.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm2.running_var"]))
1040
+ if hasattr(gps_layer.norm2, "moving_mean") and f"{p}.norm2.running_mean" in clean_dict:
1041
+ gps_layer.norm2.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm2.running_mean"]))
1042
+ if hasattr(gps_layer.norm2, "moving_variance") and f"{p}.norm2.running_var" in clean_dict:
1043
+ gps_layer.norm2.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm2.running_var"]))
1044
+
1045
+ # 4. SANGraphHead (post_mp)
1046
+ for l, fc in enumerate(model.post_mp.FC_layers):
1047
+ if f"post_mp.FC_layers.{l}.weight" in clean_dict:
1048
+ fc.kernel.assign(_to_tensor(clean_dict[f"post_mp.FC_layers.{l}.weight"].t()))
1049
+ if f"post_mp.FC_layers.{l}.bias" in clean_dict:
1050
+ fc.bias.assign(_to_tensor(clean_dict[f"post_mp.FC_layers.{l}.bias"]))
1051
+
1052
+ return model
1053
+
1054
+
1055
+ DROPBOX_CHECKPOINT_URLS = {
1056
+ "pcqm4m-GPS+RWSE.deep": "https://www.dropbox.com/s/aomimvak4gb6et3/pcqm4m-GPS%2BRWSE.deep.zip?dl=1",
1057
+ }
1058
+
1059
+
1060
+ def download_gps_checkpoint(
1061
+ checkpoint_name: str = "pcqm4m-GPS+RWSE.deep",
1062
+ cache_dir: Optional[str] = None,
1063
+ ) -> str:
1064
+ r"""Downloads and extracts a pretrained GraphGPS checkpoint.
1065
+
1066
+ Args:
1067
+ checkpoint_name (str, optional): Name of the checkpoint.
1068
+ Currently supported: ``"pcqm4m-GPS+RWSE.deep"``.
1069
+ cache_dir (str, optional): Cache directory to store downloaded checkpoint.
1070
+
1071
+ Returns:
1072
+ str: Absolute path to the extracted `.ckpt` checkpoint file.
1073
+ """
1074
+ if cache_dir is None:
1075
+ cache_dir = os.path.expanduser("~/.cache/k3_node/graphgps")
1076
+ os.makedirs(cache_dir, exist_ok=True)
1077
+
1078
+ expected_ckpt = os.path.join(cache_dir, checkpoint_name, "0", "ckpt", "148.ckpt")
1079
+ if os.path.exists(expected_ckpt):
1080
+ return expected_ckpt
1081
+
1082
+ # Also check repo pretrained dir if available
1083
+ local_repo_ckpt = os.path.join(
1084
+ os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
1085
+ "..",
1086
+ "GraphGPS",
1087
+ "pretrained",
1088
+ checkpoint_name,
1089
+ "0",
1090
+ "ckpt",
1091
+ "148.ckpt",
1092
+ )
1093
+ if os.path.exists(local_repo_ckpt):
1094
+ return os.path.abspath(local_repo_ckpt)
1095
+
1096
+ url = DROPBOX_CHECKPOINT_URLS.get(checkpoint_name)
1097
+ if url is None:
1098
+ raise ValueError(f"Unknown checkpoint '{checkpoint_name}'")
1099
+
1100
+ zip_path = os.path.join(cache_dir, f"{checkpoint_name}.zip")
1101
+ if not os.path.exists(zip_path):
1102
+ print(f"Downloading {checkpoint_name} from {url}...")
1103
+ req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0"})
1104
+ with urllib.request.urlopen(req) as resp, open(zip_path, "wb") as f:
1105
+ while True:
1106
+ chunk = resp.read(1024 * 1024)
1107
+ if not chunk:
1108
+ break
1109
+ f.write(chunk)
1110
+
1111
+ with zipfile.ZipFile(zip_path, "r") as zip_ref:
1112
+ zip_ref.extractall(cache_dir)
1113
+
1114
+ if not os.path.exists(expected_ckpt):
1115
+ # Search for any ckpt file inside extracted directory
1116
+ for root, _, files in os.walk(cache_dir):
1117
+ for file in files:
1118
+ if file.endswith(".ckpt"):
1119
+ return os.path.join(root, file)
1120
+ raise FileNotFoundError(f"Could not find .ckpt inside extracted {zip_path}")
1121
+
1122
+ return expected_ckpt