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
k3_node/models/gpse.py ADDED
@@ -0,0 +1,638 @@
1
+ from typing import Optional, List, Union
2
+ import numpy as np
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.conv import ResGatedGraphConv
7
+ from k3_node.layers.pool import global_add_pool, global_max_pool, global_mean_pool
8
+
9
+
10
+ class BatchNorm1dNode(keras.layers.Layer):
11
+ def __init__(self, channels: int, **kwargs):
12
+ super().__init__(**kwargs)
13
+ self.bn = keras.layers.BatchNormalization(epsilon=1e-5, momentum=0.9)
14
+
15
+ def build(self, input_shape=None):
16
+ self.built = True
17
+
18
+ def call(self, x, training=False):
19
+ return self.bn(x, training=training)
20
+
21
+
22
+ class BatchNorm1dEdge(keras.layers.Layer):
23
+ def __init__(self, channels: int, **kwargs):
24
+ super().__init__(**kwargs)
25
+ self.bn = keras.layers.BatchNormalization(epsilon=1e-5, momentum=0.9)
26
+
27
+ def build(self, input_shape=None):
28
+ self.built = True
29
+
30
+ def call(self, edge_attr, training=False):
31
+ return self.bn(edge_attr, training=training)
32
+
33
+
34
+ class GeneralLayer(keras.layers.Layer):
35
+ """Layer ``GeneralLayer``.
36
+
37
+ Example:
38
+ ```python
39
+ import numpy as np
40
+ from k3_node.models import GeneralLayer
41
+
42
+ x = np.random.rand(6, 16).astype("float32")
43
+ edge_index = np.array([[0, 1, 2, 3, 4, 5], [1, 2, 0, 4, 5, 3]])
44
+
45
+ layer = GeneralLayer("resgatedgcnconv", in_channels=16, out_channels=32) # conv + batch norm + activation
46
+ print(tuple(layer(x, edge_index).shape)) # (6, 32)
47
+ ```
48
+ """
49
+ def __init__(
50
+ self,
51
+ name: str,
52
+ in_channels: int,
53
+ out_channels: int,
54
+ has_batch_norm: bool = True,
55
+ has_l2_norm: bool = True,
56
+ dropout: float = 0.0,
57
+ act: Optional[str] = "relu",
58
+ **kwargs,
59
+ ):
60
+ super().__init__(**kwargs)
61
+ self.has_l2_norm = has_l2_norm
62
+ self.name_type = name.lower()
63
+
64
+ if self.name_type == "linear":
65
+ self.layer = keras.layers.Dense(out_channels, use_bias=not has_batch_norm)
66
+ elif self.name_type == "resgatedgcnconv":
67
+ self.layer = ResGatedGraphConv(in_channels, out_channels, bias=not has_batch_norm)
68
+ else:
69
+ raise ValueError(f"Unknown layer type '{name}'")
70
+
71
+ self.has_batch_norm = has_batch_norm
72
+ if has_batch_norm:
73
+ self.bn = keras.layers.BatchNormalization(epsilon=1e-5, momentum=0.9)
74
+ else:
75
+ self.bn = None
76
+
77
+ self.dropout_rate = dropout
78
+ if dropout > 0:
79
+ self.drop = keras.layers.Dropout(dropout)
80
+ else:
81
+ self.drop = None
82
+
83
+ if act is not None:
84
+ self.act = keras.activations.get(act)
85
+ else:
86
+ self.act = None
87
+
88
+ def build(self, input_shape=None):
89
+ self.built = True
90
+
91
+ def call(self, x, edge_index=None, training=False):
92
+ if self.name_type == "linear":
93
+ h = self.layer(x)
94
+ else:
95
+ h = self.layer(x, edge_index)
96
+
97
+ if self.bn is not None:
98
+ h = self.bn(h, training=training)
99
+ if self.drop is not None:
100
+ h = self.drop(h, training=training)
101
+ if self.act is not None:
102
+ h = self.act(h)
103
+ if self.has_l2_norm:
104
+ norm = ops.sqrt(ops.sum(ops.power(h, 2), axis=-1, keepdims=True)) + 1e-12
105
+ h = h / norm
106
+ return h
107
+
108
+
109
+ class GeneralMultiLayer(keras.layers.Layer):
110
+ """Layer ``GeneralMultiLayer``.
111
+
112
+ Example:
113
+ ```python
114
+ import numpy as np
115
+ from k3_node.models import GeneralMultiLayer
116
+
117
+ x = np.random.rand(6, 16).astype("float32")
118
+ edge_index = np.array([[0, 1, 2, 3, 4, 5], [1, 2, 0, 4, 5, 3]])
119
+
120
+ layer = GeneralMultiLayer("linear", in_channels=16, out_channels=32, num_layers=2)
121
+ print(tuple(layer(x).shape)) # (6, 32)
122
+ ```
123
+ """
124
+ def __init__(
125
+ self,
126
+ name: str,
127
+ in_channels: int,
128
+ out_channels: int,
129
+ hidden_channels: Optional[int] = None,
130
+ num_layers: int = 1,
131
+ has_batch_norm: bool = True,
132
+ has_l2_norm: bool = True,
133
+ dropout: float = 0.0,
134
+ act: str = "relu",
135
+ final_act: bool = True,
136
+ **kwargs,
137
+ ):
138
+ super().__init__(**kwargs)
139
+ hidden_channels = hidden_channels or out_channels
140
+ self.layers_list = []
141
+ for i in range(num_layers):
142
+ d_in = in_channels if i == 0 else hidden_channels
143
+ d_out = out_channels if i == num_layers - 1 else hidden_channels
144
+ act_i = None if i == num_layers - 1 and not final_act else act
145
+ self.layers_list.append(
146
+ GeneralLayer(
147
+ name=name,
148
+ in_channels=d_in,
149
+ out_channels=d_out,
150
+ has_batch_norm=has_batch_norm,
151
+ has_l2_norm=has_l2_norm,
152
+ dropout=dropout,
153
+ act=act_i,
154
+ )
155
+ )
156
+
157
+ def build(self, input_shape=None):
158
+ self.built = True
159
+
160
+ def call(self, x, edge_index=None, training=False):
161
+ for layer in self.layers_list:
162
+ x = layer(x, edge_index=edge_index, training=training)
163
+ return x
164
+
165
+
166
+ class GNNStackStage(keras.layers.Layer):
167
+ """Layer ``GNNStackStage``.
168
+
169
+ Example:
170
+ ```python
171
+ import numpy as np
172
+ from k3_node.models import GNNStackStage
173
+
174
+ x = np.random.rand(6, 16).astype("float32")
175
+ edge_index = np.array([[0, 1, 2, 3, 4, 5], [1, 2, 0, 4, 5, 3]])
176
+
177
+ stage = GNNStackStage(in_channels=16, out_channels=16, num_layers=2) # stack of message-passing layers
178
+ print(tuple(stage(x, edge_index).shape)) # (6, 16)
179
+ ```
180
+ """
181
+ def __init__(
182
+ self,
183
+ in_channels: int,
184
+ out_channels: int,
185
+ num_layers: int,
186
+ layer_type: str = "resgatedgcnconv",
187
+ stage_type: str = "skipsum",
188
+ final_l2_norm: bool = True,
189
+ has_batch_norm: bool = True,
190
+ has_l2_norm: bool = True,
191
+ dropout: float = 0.2,
192
+ act: Optional[str] = "relu",
193
+ **kwargs,
194
+ ):
195
+ super().__init__(**kwargs)
196
+ self.num_layers = num_layers
197
+ self.stage_type = stage_type
198
+ self.final_l2_norm = final_l2_norm
199
+
200
+ self.layers_list = []
201
+ for i in range(num_layers):
202
+ if stage_type == "skipconcat":
203
+ d_in = in_channels if i == 0 else in_channels + i * out_channels
204
+ else:
205
+ d_in = in_channels if i == 0 else out_channels
206
+ self.layers_list.append(
207
+ GeneralLayer(
208
+ name=layer_type,
209
+ in_channels=d_in,
210
+ out_channels=out_channels,
211
+ has_batch_norm=has_batch_norm,
212
+ has_l2_norm=has_l2_norm,
213
+ dropout=dropout,
214
+ act=act,
215
+ )
216
+ )
217
+
218
+ def call(self, x, edge_index, training=False):
219
+ for i, layer in enumerate(self.layers_list):
220
+ prev_x = x
221
+ h = layer(x, edge_index=edge_index, training=training)
222
+ if self.stage_type == "skipsum":
223
+ x = prev_x + h
224
+ elif self.stage_type == "skipconcat" and i < self.num_layers - 1:
225
+ x = ops.concatenate([prev_x, h], axis=1)
226
+ else:
227
+ x = h
228
+
229
+ if self.final_l2_norm:
230
+ norm = ops.sqrt(ops.sum(ops.power(x, 2), axis=-1, keepdims=True)) + 1e-12
231
+ x = x / norm
232
+
233
+ return x
234
+
235
+
236
+ class IdentityHead(keras.layers.Layer):
237
+ def call(self, x, **kwargs):
238
+ return x
239
+
240
+
241
+ class GNNInductiveHybridMultiHead(keras.layers.Layer):
242
+ r"""GNN prediction head for inductive node and graph prediction tasks.
243
+
244
+ Example:
245
+ ```python
246
+ import numpy as np
247
+ from k3_node.models import GNNInductiveHybridMultiHead
248
+
249
+ x = np.random.rand(6, 16).astype("float32")
250
+ edge_index = np.array([[0, 1, 2, 3, 4, 5], [1, 2, 0, 4, 5, 3]])
251
+ batch = np.array([0, 0, 0, 1, 1, 1]) # two graphs
252
+
253
+ head = GNNInductiveHybridMultiHead(dim_in=16, dim_out=4, num_node_targets=3, num_graph_targets=2,
254
+ layers_post_mp=1)
255
+ node_pred, graph_pred = head(x, batch=batch, batch_size=2)
256
+ print(tuple(node_pred.shape), tuple(graph_pred.shape)) # (6, 3) (2, 2)
257
+ ```
258
+ """
259
+ def __init__(
260
+ self,
261
+ dim_in: int,
262
+ dim_out: int,
263
+ num_node_targets: int,
264
+ num_graph_targets: int,
265
+ layers_post_mp: int,
266
+ virtual_node: bool = True,
267
+ multi_head_dim_inner: int = 32,
268
+ graph_pooling: str = "add",
269
+ has_bn: bool = True,
270
+ has_l2norm: bool = True,
271
+ dropout: float = 0.2,
272
+ act: str = "relu",
273
+ **kwargs,
274
+ ):
275
+ super().__init__(**kwargs)
276
+ self.node_target_dim = num_node_targets
277
+ self.graph_target_dim = num_graph_targets
278
+ self.virtual_node = virtual_node
279
+ self.graph_pooling = graph_pooling
280
+
281
+ self.node_post_mps = [
282
+ GeneralMultiLayer(
283
+ name="linear",
284
+ in_channels=dim_in,
285
+ out_channels=1,
286
+ hidden_channels=multi_head_dim_inner,
287
+ num_layers=layers_post_mp,
288
+ has_batch_norm=has_bn,
289
+ has_l2_norm=has_l2norm,
290
+ dropout=dropout,
291
+ act=act,
292
+ final_act=False,
293
+ )
294
+ for _ in range(num_node_targets)
295
+ ]
296
+
297
+ self.graph_post_mp = GeneralMultiLayer(
298
+ name="linear",
299
+ in_channels=dim_in,
300
+ out_channels=num_graph_targets,
301
+ hidden_channels=dim_in,
302
+ num_layers=layers_post_mp,
303
+ has_batch_norm=has_bn,
304
+ has_l2_norm=has_l2norm,
305
+ dropout=dropout,
306
+ act=act,
307
+ final_act=False,
308
+ )
309
+
310
+ def call(self, x, batch=None, training=False, batch_size=None):
311
+ node_feats = [m(x, training=training) for m in self.node_post_mps]
312
+ node_pred = ops.concatenate(node_feats, axis=-1)
313
+
314
+ if batch is None:
315
+ batch = ops.zeros(ops.shape(x)[:1], dtype="int32")
316
+ batch_size = 1
317
+ else:
318
+ batch = ops.cast(batch, "int32")
319
+
320
+ if self.graph_pooling == "max":
321
+ graph_emb = global_max_pool(x, batch, size=batch_size)
322
+ elif self.graph_pooling == "mean":
323
+ graph_emb = global_mean_pool(x, batch, size=batch_size)
324
+ else:
325
+ graph_emb = global_add_pool(x, batch, size=batch_size)
326
+
327
+ graph_pred = self.graph_post_mp(graph_emb, training=training)
328
+ return node_pred, graph_pred
329
+
330
+
331
+ class GPSE(keras.layers.Layer):
332
+ r"""The Graph Positional and Structural Encoder (GPSE) model from the
333
+ `"Graph Positional and Structural Encoder"
334
+ <https://arxiv.org/abs/2307.07107>`_ paper.
335
+
336
+ Example:
337
+ ```python
338
+ import numpy as np
339
+ from k3_node.models import GPSE
340
+
341
+ x = np.random.rand(6, 16).astype("float32")
342
+ edge_index = np.array([[0, 1, 2, 3, 4, 5], [1, 2, 0, 4, 5, 3]])
343
+
344
+ model = GPSE(dim_in=16, dim_inner=32, layers_pre_mp=1, layers_mp=2, layers_post_mp=1, use_repr=True)
345
+ print(tuple(model(x, edge_index).shape)) # (6, 32): positional/structural encodings per node
346
+ ```
347
+ """
348
+ def __init__(
349
+ self,
350
+ dim_in: int = 20,
351
+ dim_out: int = 51,
352
+ dim_inner: int = 512,
353
+ layer_type: str = "resgatedgcnconv",
354
+ layers_pre_mp: int = 1,
355
+ layers_mp: int = 20,
356
+ layers_post_mp: int = 2,
357
+ num_node_targets: int = 51,
358
+ num_graph_targets: int = 11,
359
+ stage_type: str = "skipsum",
360
+ has_bn: bool = True,
361
+ head_bn: bool = False,
362
+ final_l2norm: bool = True,
363
+ has_l2norm: bool = True,
364
+ dropout: float = 0.2,
365
+ has_act: bool = True,
366
+ final_act: bool = True,
367
+ act: str = "relu",
368
+ virtual_node: bool = True,
369
+ multi_head_dim_inner: int = 32,
370
+ graph_pooling: str = "add",
371
+ use_repr: bool = True,
372
+ repr_type: str = "no_post_mp",
373
+ bernoulli_threshold: float = 0.5,
374
+ **kwargs,
375
+ ):
376
+ super().__init__(**kwargs)
377
+ self.use_repr = use_repr
378
+ self.repr_type = repr_type
379
+ self.dim_inner = dim_inner
380
+ self.bernoulli_threshold = bernoulli_threshold
381
+
382
+ if layers_pre_mp > 0:
383
+ self.pre_mp = GeneralMultiLayer(
384
+ name="linear",
385
+ in_channels=dim_in,
386
+ out_channels=dim_inner,
387
+ hidden_channels=dim_inner,
388
+ num_layers=layers_pre_mp,
389
+ has_batch_norm=has_bn,
390
+ has_l2_norm=has_l2norm,
391
+ dropout=dropout,
392
+ act=act,
393
+ final_act=final_act,
394
+ )
395
+ d_in = dim_inner
396
+ else:
397
+ self.pre_mp = None
398
+ d_in = dim_in
399
+
400
+ if layers_mp > 0:
401
+ self.mp = GNNStackStage(
402
+ in_channels=d_in,
403
+ out_channels=dim_inner,
404
+ num_layers=layers_mp,
405
+ layer_type=layer_type,
406
+ stage_type=stage_type,
407
+ final_l2_norm=final_l2norm,
408
+ has_batch_norm=has_bn,
409
+ has_l2_norm=has_l2norm,
410
+ dropout=dropout,
411
+ act=act if has_act else None,
412
+ )
413
+ else:
414
+ self.mp = None
415
+
416
+ if use_repr:
417
+ self.post_mp = IdentityHead()
418
+ else:
419
+ self.post_mp = GNNInductiveHybridMultiHead(
420
+ dim_in=dim_inner,
421
+ dim_out=dim_out,
422
+ num_node_targets=num_node_targets,
423
+ num_graph_targets=num_graph_targets,
424
+ layers_post_mp=layers_post_mp,
425
+ virtual_node=virtual_node,
426
+ multi_head_dim_inner=multi_head_dim_inner,
427
+ graph_pooling=graph_pooling,
428
+ has_bn=head_bn,
429
+ has_l2norm=has_l2norm,
430
+ dropout=dropout,
431
+ act=act,
432
+ )
433
+
434
+ def build(self, input_shape=None):
435
+ self.built = True
436
+
437
+ def call(self, x, edge_index=None, batch=None, training=False, batch_size=None):
438
+ # Support both (x, edge_index, batch) and batch object with attributes
439
+ if hasattr(x, "x") and hasattr(x, "edge_index"):
440
+ edge_index = x.edge_index
441
+ batch = getattr(x, "batch", None)
442
+ x = x.x
443
+
444
+ if self.pre_mp is not None:
445
+ x = self.pre_mp(x, training=training)
446
+
447
+ if self.mp is not None:
448
+ x = self.mp(x, edge_index, training=training)
449
+
450
+ if self.use_repr:
451
+ return x
452
+
453
+ return self.post_mp(x, batch=batch, training=training, batch_size=batch_size)
454
+
455
+
456
+ class GPSENodeEncoder(keras.layers.Layer):
457
+ r"""A helper linear/MLP encoder that takes the :class:`GPSE` encodings
458
+ precomputed in the input graphs, maps them to a desired
459
+ dimension defined by :obj:`dim_pe_out` and appends them to node features.
460
+
461
+ Args:
462
+ dim_emb (int): Size of final node embedding.
463
+ dim_pe_in (int): Original dimension of GPSE encodings.
464
+ dim_pe_out (int): Desired dimension of GPSE encodings after the encoder.
465
+ dim_in (int, optional): Original dimension of input node features. (default: None)
466
+ expand_x (bool, optional): Expand node features x. (default: False)
467
+ norm_type (str, optional): Type of normalization. (default: "batchnorm")
468
+ model_type (str, optional): Encoder model ('mlp' or 'linear'). (default: "mlp")
469
+ n_layers (int, optional): Number of MLP layers. (default: 2)
470
+ dropout_be (float, optional): Dropout before encoding. (default: 0.5)
471
+ dropout_ae (float, optional): Dropout after encoding. (default: 0.2)
472
+
473
+ Example:
474
+ ```python
475
+ import numpy as np
476
+ from k3_node.models import GPSENodeEncoder
477
+
478
+ x = np.random.rand(6, 16).astype("float32")
479
+ edge_index = np.array([[0, 1, 2, 3, 4, 5], [1, 2, 0, 4, 5, 3]])
480
+ pos_enc = np.random.rand(6, 32).astype("float32") # encodings produced by GPSE
481
+
482
+ encoder = GPSENodeEncoder(dim_emb=64, dim_pe_in=32, dim_pe_out=16, dim_in=16, expand_x=True)
483
+ print(tuple(encoder(x, pos_enc).shape)) # (6, 64): node features concatenated with projected encodings
484
+ ```
485
+ """
486
+ def __init__(
487
+ self,
488
+ dim_emb: int,
489
+ dim_pe_in: int,
490
+ dim_pe_out: int,
491
+ dim_in: Optional[int] = None,
492
+ expand_x: bool = False,
493
+ norm_type: str = "batchnorm",
494
+ model_type: str = "mlp",
495
+ n_layers: int = 2,
496
+ dropout_be: float = 0.5,
497
+ dropout_ae: float = 0.2,
498
+ **kwargs,
499
+ ):
500
+ super().__init__(**kwargs)
501
+ assert dim_emb > dim_pe_out, (
502
+ "Desired GPSE dimension (dim_pe_out) must be smaller than "
503
+ "the final node embedding dimension (dim_emb)."
504
+ )
505
+
506
+ self.expand_x = expand_x
507
+ if expand_x:
508
+ self.linear_x = keras.layers.Dense(dim_emb - dim_pe_out)
509
+ else:
510
+ self.linear_x = None
511
+
512
+ if norm_type == "batchnorm":
513
+ self.raw_norm = keras.layers.BatchNormalization(momentum=0.9, epsilon=1e-5)
514
+ else:
515
+ self.raw_norm = None
516
+
517
+ self.dropout_be = keras.layers.Dropout(dropout_be)
518
+ self.dropout_ae = keras.layers.Dropout(dropout_ae)
519
+
520
+ if model_type == "mlp":
521
+ layers = []
522
+ if n_layers == 1:
523
+ layers.append(keras.layers.Dense(dim_pe_out, activation="relu"))
524
+ else:
525
+ layers.append(keras.layers.Dense(2 * dim_pe_out, activation="relu"))
526
+ for _ in range(n_layers - 2):
527
+ layers.append(keras.layers.Dense(2 * dim_pe_out, activation="relu"))
528
+ layers.append(keras.layers.Dense(dim_pe_out, activation="relu"))
529
+ self.pe_encoder = keras.Sequential(layers)
530
+ elif model_type == "linear":
531
+ self.pe_encoder = keras.layers.Dense(dim_pe_out)
532
+ else:
533
+ raise ValueError(f"Does not support '{model_type}' encoder model.")
534
+
535
+ def call(self, x, pos_enc, training=False):
536
+ pos_enc = self.dropout_be(pos_enc, training=training)
537
+ if self.raw_norm is not None:
538
+ pos_enc = self.raw_norm(pos_enc, training=training)
539
+ pos_enc = self.pe_encoder(pos_enc)
540
+ pos_enc = self.dropout_ae(pos_enc, training=training)
541
+
542
+ h = self.linear_x(x) if self.expand_x else x
543
+ return ops.concatenate([h, pos_enc], axis=1)
544
+
545
+
546
+
547
+ GPSE_URLS = {
548
+ 'molpcba': 'https://zenodo.org/record/8145095/files/gpse_model_molpcba_1.0.pt',
549
+ 'zinc': 'https://zenodo.org/record/8145095/files/gpse_model_zinc_1.0.pt',
550
+ 'pcqm4mv2': 'https://zenodo.org/record/8145095/files/gpse_model_pcqm4mv2_1.0.pt',
551
+ 'geom': 'https://zenodo.org/record/8145095/files/gpse_model_geom_1.0.pt',
552
+ 'chembl': 'https://zenodo.org/record/8145095/files/gpse_model_chembl_1.0.pt',
553
+ }
554
+
555
+
556
+ def _gpse_from_pretrained(cls, name: str, root: str = 'GPSE_pretrained'):
557
+ r"""Returns a :class:`GPSE` pre-trained on ``name`` (``"molpcba"``, ``"zinc"``, ``"pcqm4mv2"``,
558
+ ``"geom"`` or ``"chembl"``), with the authors' weights converted from PyTorch (reading them
559
+ needs PyTorch). The model returns the 512-dimensional node representations."""
560
+ import os
561
+ import os.path as osp
562
+
563
+ import torch
564
+
565
+ from k3_node.data.download import download_url
566
+
567
+ root = osp.expanduser(root)
568
+ os.makedirs(root, exist_ok=True)
569
+ path = osp.join(root, GPSE_URLS[name].rsplit('/', 1)[1])
570
+ if not osp.exists(path):
571
+ path = download_url(GPSE_URLS[name], root)
572
+ state = torch.load(path, map_location='cpu', weights_only=False)['model_state']
573
+ state = {k.split('.', 1)[1]: v.detach().cpu().numpy() for k, v in state.items()}
574
+
575
+ model = cls() # all pre-trained models use the default arguments
576
+ model(np.zeros((3, 20), 'float32'), np.array([[0, 1, 2], [1, 2, 0]])) # create the weights
577
+
578
+ def batch_norm(bn, prefix):
579
+ bn.gamma.assign(state[f'{prefix}.weight'])
580
+ bn.beta.assign(state[f'{prefix}.bias'])
581
+ bn.moving_mean.assign(state[f'{prefix}.running_mean'])
582
+ bn.moving_variance.assign(state[f'{prefix}.running_var'])
583
+
584
+ for i, layer in enumerate(model.pre_mp.layers_list):
585
+ layer.layer.kernel.assign(state[f'pre_mp.Layer_{i}.layer.model.weight'].T)
586
+ batch_norm(layer.bn, f'pre_mp.Layer_{i}.post_layer.0')
587
+ for i, layer in enumerate(model.mp.layers_list):
588
+ prefix = f'mp.layer{i}.layer.model'
589
+ for lin in ['lin_key', 'lin_query', 'lin_value']:
590
+ getattr(layer.layer, lin).kernel.assign(state[f'{prefix}.{lin}.weight'].T)
591
+ getattr(layer.layer, lin).bias.assign(state[f'{prefix}.{lin}.bias'])
592
+ layer.layer.lin_skip.kernel.assign(state[f'{prefix}.lin_skip.weight'].T)
593
+ batch_norm(layer.bn, f'mp.layer{i}.post_layer.0')
594
+ return model
595
+
596
+
597
+ GPSE.from_pretrained = classmethod(_gpse_from_pretrained)
598
+
599
+
600
+ def _random_features(n, dim, rand_type, bernoulli_threshold=0.5):
601
+ if rand_type == 'NormalSE':
602
+ return np.random.normal(0.0, 1.0, size=(n, dim)).astype('float32')
603
+ if rand_type == 'UniformSE':
604
+ return np.random.uniform(0.0, 1.0, size=(n, dim)).astype('float32')
605
+ if rand_type == 'BernoulliSE':
606
+ return (np.random.uniform(0.0, 1.0, size=(n, dim)) < bernoulli_threshold).astype('float32')
607
+ raise ValueError(f'Unknown rand_type {rand_type!r}')
608
+
609
+
610
+ def gpse_encodings(model, graphs, use_vn: bool = True, rand_type: str = 'NormalSE', dim_in: int = 20):
611
+ r"""Computes the GPSE encodings of a list of graphs at once: the model reads random node
612
+ features (plus a virtual node connected to all nodes if ``use_vn``) and returns one
613
+ representation per node. Returns one ``[num_nodes, dim]`` array per graph."""
614
+ from k3_node.data import Batch, Data
615
+ from k3_node.transforms import VirtualNode
616
+
617
+ structures = []
618
+ for g in graphs:
619
+ s = Data(edge_index=np.asarray(ops.convert_to_numpy(g.edge_index)), num_nodes=g.num_nodes)
620
+ structures.append(VirtualNode()(s) if use_vn else s)
621
+ batch = Batch.from_data_list(structures)
622
+ ptr = np.asarray(ops.convert_to_numpy(batch.ptr))
623
+ x = _random_features(int(ptr[-1]), dim_in, rand_type)
624
+ if use_vn:
625
+ x[ptr[1:] - 1] = 0 # the virtual nodes start from zero
626
+ out = np.asarray(ops.convert_to_numpy(model(x, batch.edge_index, training=False)))
627
+ return [out[start:end - (1 if use_vn else 0)] for start, end in zip(ptr[:-1], ptr[1:])]
628
+
629
+
630
+ def precompute_gpse(model, dataset, use_vn: bool = True, rand_type: str = 'NormalSE', batch_size: int = 128):
631
+ r"""Returns the graphs of ``dataset`` as a list, each with its GPSE encodings in
632
+ ``pestat_GPSE`` (as PyG's ``precompute_GPSE``)."""
633
+ graphs = [dataset[i] for i in range(len(dataset))]
634
+ for start in range(0, len(graphs), batch_size):
635
+ chunk = graphs[start:start + batch_size]
636
+ for graph, enc in zip(chunk, gpse_encodings(model, chunk, use_vn, rand_type)):
637
+ graph.pestat_GPSE = enc
638
+ return graphs