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,108 @@
1
+ import keras
2
+ from keras import ops
3
+
4
+
5
+ class ARLinkPredictor(keras.layers.Layer):
6
+ r"""Link predictor using Attract-Repel embeddings from the paper
7
+ `"Pseudo-Euclidean Attract-Repel Embeddings for Undirected Graphs"
8
+ <https://arxiv.org/abs/2106.09671>`_.
9
+
10
+ This model splits node embeddings into: attract and repel.
11
+ The edge prediction score is computed as the dot product of attract
12
+ components minus the dot product of repel components.
13
+
14
+ Args:
15
+ in_channels (int): Size of each input sample.
16
+ hidden_channels (int): Size of hidden embeddings.
17
+ out_channels (int, optional): Size of output embeddings. If set to
18
+ `None`, will default to `hidden_channels`. (default: `None`)
19
+ num_layers (int): Number of message passing layers. (default: `2`)
20
+ dropout (float): Dropout probability. (default: `0.0`)
21
+ attract_ratio (float): Ratio to use for attract component. Must be
22
+ between 0 and 1. (default: `0.5`)
23
+
24
+ Example:
25
+ ```python
26
+ import numpy as np
27
+ from k3_node.models import ARLinkPredictor
28
+
29
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
30
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
31
+
32
+ model = ARLinkPredictor(in_channels=8, hidden_channels=16, num_layers=2)
33
+ scores = model(x, edge_index) # one link score per edge in edge_index
34
+ print(tuple(scores.shape)) # (30,)
35
+ ```
36
+ """
37
+ def __init__(self, in_channels, hidden_channels, out_channels=None,
38
+ num_layers=2, dropout=0.0, attract_ratio=0.5, **kwargs):
39
+ super().__init__(**kwargs)
40
+
41
+ if out_channels is None:
42
+ out_channels = hidden_channels
43
+
44
+ self.in_channels = in_channels
45
+ self.hidden_channels = hidden_channels
46
+ self.out_channels = out_channels
47
+ self.num_layers = num_layers
48
+ self.dropout_rate = dropout
49
+
50
+ if not 0 <= attract_ratio <= 1:
51
+ raise ValueError(f"attract_ratio must be between 0 and 1, got {attract_ratio}")
52
+
53
+ self.attract_ratio = attract_ratio
54
+ self.attract_dim = int(out_channels * attract_ratio)
55
+ self.repel_dim = out_channels - self.attract_dim
56
+
57
+ self.lins = [keras.layers.Dense(hidden_channels)]
58
+ for _ in range(num_layers - 2):
59
+ self.lins.append(keras.layers.Dense(hidden_channels))
60
+
61
+ self.lin_attract = keras.layers.Dense(self.attract_dim)
62
+ self.lin_repel = keras.layers.Dense(self.repel_dim)
63
+ self.dropout = keras.layers.Dropout(dropout) if dropout > 0.0 else None
64
+
65
+ self.lins[0].build((None, in_channels))
66
+ for lin in self.lins[1:]:
67
+ lin.build((None, hidden_channels))
68
+ self.lin_attract.build((None, hidden_channels))
69
+ self.lin_repel.build((None, hidden_channels))
70
+
71
+ def encode(self, x, *args, training=None, **kwargs):
72
+ r"""Encode node features into attract-repel embeddings."""
73
+ for lin in self.lins:
74
+ x = lin(x)
75
+ x = ops.relu(x)
76
+ if self.dropout is not None:
77
+ x = self.dropout(x, training=training)
78
+
79
+ attract_x = self.lin_attract(x)
80
+ repel_x = self.lin_repel(x)
81
+
82
+ return attract_x, repel_x
83
+
84
+ def decode(self, attract_z, repel_z, edge_index):
85
+ r"""Decode edge scores from attract-repel embeddings."""
86
+ row, col = edge_index[0], edge_index[1]
87
+ attract_z_row = ops.take(attract_z, row, axis=0)
88
+ attract_z_col = ops.take(attract_z, col, axis=0)
89
+ repel_z_row = ops.take(repel_z, row, axis=0)
90
+ repel_z_col = ops.take(repel_z, col, axis=0)
91
+
92
+ attract_score = ops.sum(attract_z_row * attract_z_col, axis=1)
93
+ repel_score = ops.sum(repel_z_row * repel_z_col, axis=1)
94
+
95
+ return attract_score - repel_score
96
+
97
+ def call(self, x, edge_index, training=None):
98
+ attract_z, repel_z = self.encode(x, training=training)
99
+ return ops.sigmoid(self.decode(attract_z, repel_z, edge_index))
100
+
101
+ def calculate_r_fraction(self, attract_z, repel_z):
102
+ r"""Calculate the R-fraction (proportion of energy in repel space)."""
103
+ attract_norm_squared = ops.sum(ops.square(attract_z))
104
+ repel_norm_squared = ops.sum(ops.square(repel_z))
105
+
106
+ r_fraction = repel_norm_squared / (attract_norm_squared + repel_norm_squared + 1e-10)
107
+
108
+ return float(ops.convert_to_numpy(r_fraction))
@@ -0,0 +1,318 @@
1
+ import keras
2
+ from keras import ops
3
+
4
+ from k3_node.models.utils import negative_sampling, reset
5
+
6
+ EPS = 1e-15
7
+ MAX_LOGSTD = 10
8
+
9
+
10
+ def _randn_like(x):
11
+ return keras.random.normal(ops.shape(x), dtype=x.dtype)
12
+
13
+
14
+ class InnerProductDecoder:
15
+ r"""The inner product decoder from the `"Variational Graph Auto-Encoders"
16
+ <https://arxiv.org/abs/1611.07308>`_ paper.
17
+
18
+ .. math::
19
+ \sigma(\mathbf{Z}\mathbf{Z}^{\top})
20
+
21
+ where :math:`\mathbf{Z} \in \mathbb{R}^{N \times d}` denotes the latent
22
+ space produced by the encoder.
23
+ """
24
+ def __call__(self, z, edge_index, sigmoid: bool = True):
25
+ r"""Decodes the latent variables `z` into edge probabilities for
26
+ the given node-pairs `edge_index`."""
27
+ row, col = edge_index[0], edge_index[1]
28
+ value = ops.sum(ops.take(z, row, axis=0) * ops.take(z, col, axis=0), axis=1)
29
+ return ops.sigmoid(value) if sigmoid else value
30
+
31
+ def forward_all(self, z, sigmoid: bool = True):
32
+ r"""Decodes the latent variables `z` into a probabilistic dense
33
+ adjacency matrix."""
34
+ adj = ops.matmul(z, ops.transpose(z))
35
+ return ops.sigmoid(adj) if sigmoid else adj
36
+
37
+
38
+ class GAE:
39
+ r"""The Graph Auto-Encoder model from the
40
+ `"Variational Graph Auto-Encoders" <https://arxiv.org/abs/1611.07308>`_
41
+ paper based on user-defined encoder and decoder models.
42
+
43
+ Args:
44
+ encoder: The encoder module.
45
+ decoder (optional): The decoder module. If set to `None`, will
46
+ default to `InnerProductDecoder`. (default: `None`)
47
+ """
48
+ def __init__(self, encoder, decoder=None):
49
+ self.encoder = encoder
50
+ self.decoder = InnerProductDecoder() if decoder is None else decoder
51
+ self.reset_parameters()
52
+
53
+ def reset_parameters(self):
54
+ r"""Resets all learnable parameters of the module."""
55
+ reset(self.encoder)
56
+ reset(self.decoder)
57
+
58
+ def __call__(self, *args, **kwargs):
59
+ r"""Alias for `encode`."""
60
+ return self.encoder(*args, **kwargs)
61
+
62
+ def encode(self, *args, **kwargs):
63
+ r"""Runs the encoder and computes node-wise latent variables."""
64
+ return self.encoder(*args, **kwargs)
65
+
66
+ def decode(self, *args, **kwargs):
67
+ r"""Runs the decoder and computes edge probabilities."""
68
+ return self.decoder(*args, **kwargs)
69
+
70
+ def eval(self):
71
+ self.training = False
72
+ return self
73
+
74
+ def train(self):
75
+ self.training = True
76
+ return self
77
+
78
+ def to(self, *args, **kwargs):
79
+ return self
80
+
81
+ def recon_loss(self, z, pos_edge_index, neg_edge_index=None):
82
+ r"""Given latent variables `z`, computes the binary cross entropy
83
+ loss for positive edges `pos_edge_index` and negative sampled
84
+ edges."""
85
+ pos_loss = -ops.mean(ops.log(self.decoder(z, pos_edge_index, sigmoid=True) + EPS))
86
+
87
+ if neg_edge_index is None:
88
+ neg_edge_index = negative_sampling(pos_edge_index, ops.shape(z)[0])
89
+ neg_loss = -ops.mean(ops.log(1 - self.decoder(z, neg_edge_index, sigmoid=True) + EPS))
90
+
91
+ return pos_loss + neg_loss
92
+
93
+ def test(self, z, pos_edge_index, neg_edge_index):
94
+ r"""Given latent variables `z`, positive edges `pos_edge_index` and
95
+ negative edges `neg_edge_index`, computes area under the ROC curve
96
+ (AUC) and average precision (AP) scores."""
97
+ from sklearn.metrics import average_precision_score, roc_auc_score
98
+
99
+ pos_y = ops.ones((ops.shape(pos_edge_index)[1],))
100
+ neg_y = ops.zeros((ops.shape(neg_edge_index)[1],))
101
+ y = ops.concatenate([pos_y, neg_y], axis=0)
102
+
103
+ pos_pred = self.decoder(z, pos_edge_index, sigmoid=True)
104
+ neg_pred = self.decoder(z, neg_edge_index, sigmoid=True)
105
+ pred = ops.concatenate([pos_pred, neg_pred], axis=0)
106
+
107
+ y, pred = ops.convert_to_numpy(y), ops.convert_to_numpy(pred)
108
+
109
+ return roc_auc_score(y, pred), average_precision_score(y, pred)
110
+
111
+ # ---- Keras-style training -------------------------------------------------------------------
112
+ def compile(self, optimizer, discriminator_optimizer=None, discriminator_steps: int = 5):
113
+ r"""Sets the optimizers used by :meth:`fit`.
114
+
115
+ Args:
116
+ optimizer (keras.optimizers.Optimizer): Trains the encoder (and decoder).
117
+ discriminator_optimizer (keras.optimizers.Optimizer, optional): Trains the
118
+ discriminator of adversarial models (:class:`ARGA`, :class:`ARGVA`).
119
+ discriminator_steps (int): Discriminator updates per encoder update. (default: ``5``)
120
+ """
121
+ self.optimizer = optimizer
122
+ self.discriminator_optimizer = discriminator_optimizer
123
+ self.discriminator_steps = discriminator_steps
124
+
125
+ def _encoder_loss(self, data):
126
+ z = self.encode(data.x, data.edge_index, training=True)
127
+ loss = self.recon_loss(z, data.pos_edge_label_index)
128
+ if isinstance(self, ARGA):
129
+ loss = loss + self.reg_loss(z)
130
+ if hasattr(self, "kl_loss"):
131
+ loss = loss + (1 / data.num_nodes) * self.kl_loss()
132
+ return loss
133
+
134
+ def train_step(self, data):
135
+ r"""Runs one training step on ``data`` and returns the loss."""
136
+ from k3_node.training import gradient_step
137
+
138
+ self.train()
139
+ if isinstance(self, ARGA):
140
+ z = ops.stop_gradient(self.encode(data.x, data.edge_index, training=True))
141
+ for _ in range(self.discriminator_steps):
142
+ gradient_step(lambda: self.discriminator_loss(z), self.discriminator.trainable_variables,
143
+ self.discriminator_optimizer)
144
+ return gradient_step(lambda: self._encoder_loss(data), self._trainable_variables(), self.optimizer)
145
+
146
+ def _trainable_variables(self):
147
+ variables = list(self.encoder.trainable_variables)
148
+ return variables + list(getattr(self.decoder, "trainable_variables", []))
149
+
150
+ def fit(self, data, epochs: int = 1, validation_data=None, verbose: int = 1):
151
+ r"""Trains the model on one graph for ``epochs`` full-graph steps.
152
+
153
+ Args:
154
+ data (Data): The training graph with node features ``x``, the message passing edges
155
+ ``edge_index`` and the edges to reconstruct ``pos_edge_label_index``, as created
156
+ by :class:`~k3_node.transforms.RandomLinkSplit` with ``split_labels=True``.
157
+ epochs (int): The number of training steps. (default: ``1``)
158
+ validation_data (Data, optional): A graph with ``pos_edge_label_index`` and
159
+ ``neg_edge_label_index`` on which AUC and average precision are reported.
160
+ verbose (int): ``0`` is silent, otherwise one line is printed per epoch.
161
+
162
+ Returns:
163
+ dict: The loss (and validation metrics) of every epoch.
164
+ """
165
+ if getattr(self, "optimizer", None) is None:
166
+ raise ValueError("Call `compile(optimizer=...)` before `fit`.")
167
+ if isinstance(self, ARGA) and self.discriminator_optimizer is None:
168
+ raise ValueError("Adversarial models need `compile(..., discriminator_optimizer=...)`.")
169
+ # Create the variables before the first gradient step
170
+ z = self.encode(data.x, data.edge_index)
171
+ if isinstance(self, ARGA):
172
+ self.discriminator(z)
173
+
174
+ history = {"loss": []}
175
+ for epoch in range(1, epochs + 1):
176
+ logs = {"loss": self.train_step(data)}
177
+ if validation_data is not None:
178
+ logs.update({f"val_{k}": v for k, v in self.evaluate(validation_data).items()})
179
+ for key, value in logs.items():
180
+ history.setdefault(key, []).append(value)
181
+ if verbose:
182
+ print(f"Epoch {epoch:03d}: " + ", ".join(f"{k}: {v:.4f}" for k, v in logs.items()))
183
+ return history
184
+
185
+ def evaluate(self, data):
186
+ r"""Returns the link prediction AUC and average precision on ``data`` (which needs
187
+ ``pos_edge_label_index`` and ``neg_edge_label_index``)."""
188
+ z = self.embed(data)
189
+ auc, ap = self.test(z, data.pos_edge_label_index, data.neg_edge_label_index)
190
+ return {"auc": float(auc), "ap": float(ap)}
191
+
192
+ def embed(self, data):
193
+ r"""Returns the node embeddings of ``data`` (without sampling noise)."""
194
+ from k3_node.training import no_grad
195
+
196
+ self.eval()
197
+ with no_grad():
198
+ return self.encode(data.x, data.edge_index, training=False)
199
+
200
+
201
+ class VGAE(GAE):
202
+ r"""The Variational Graph Auto-Encoder model from the
203
+ `"Variational Graph Auto-Encoders" <https://arxiv.org/abs/1611.07308>`_
204
+ paper.
205
+
206
+ Args:
207
+ encoder: The encoder module to compute :math:`\mu` and
208
+ :math:`\log\sigma^2`.
209
+ decoder (optional): The decoder module. If set to `None`, will
210
+ default to `InnerProductDecoder`. (default: `None`)
211
+ """
212
+ def __init__(self, encoder, decoder=None):
213
+ super().__init__(encoder, decoder)
214
+ self.training = True
215
+
216
+ def reparametrize(self, mu, logstd):
217
+ if self.training:
218
+ return mu + _randn_like(logstd) * ops.exp(logstd)
219
+ return mu
220
+
221
+ def encode(self, *args, **kwargs):
222
+ self._mu, self._logstd = self.encoder(*args, **kwargs)
223
+ self._logstd = ops.minimum(self._logstd, MAX_LOGSTD)
224
+ z = self.reparametrize(self._mu, self._logstd)
225
+ return z
226
+
227
+ def kl_loss(self, mu=None, logstd=None):
228
+ r"""Computes the KL loss, either for the passed arguments `mu` and
229
+ `logstd`, or based on latent variables from last encoding."""
230
+ mu = self._mu if mu is None else mu
231
+ logstd = self._logstd if logstd is None else ops.minimum(logstd, MAX_LOGSTD)
232
+ return -0.5 * ops.mean(
233
+ ops.sum(1 + 2 * logstd - ops.square(mu) - ops.square(ops.exp(logstd)), axis=1)
234
+ )
235
+
236
+ def eval(self):
237
+ self.training = False
238
+
239
+ def train(self):
240
+ self.training = True
241
+
242
+
243
+ class ARGA(GAE):
244
+ r"""The Adversarially Regularized Graph Auto-Encoder model from the
245
+ `"Adversarially Regularized Graph Autoencoder for Graph Embedding"
246
+ <https://arxiv.org/abs/1802.04407>`_ paper.
247
+
248
+ Args:
249
+ encoder: The encoder module.
250
+ discriminator: The discriminator module.
251
+ decoder (optional): The decoder module. If set to `None`, will
252
+ default to `InnerProductDecoder`. (default: `None`)
253
+ """
254
+ def __init__(self, encoder, discriminator, decoder=None):
255
+ super().__init__(encoder, decoder)
256
+ self.discriminator = discriminator
257
+ reset(self.discriminator)
258
+
259
+ def reset_parameters(self):
260
+ super().reset_parameters()
261
+ reset(getattr(self, "discriminator", None))
262
+
263
+ def reg_loss(self, z):
264
+ r"""Computes the regularization loss of the encoder."""
265
+ real = ops.sigmoid(self.discriminator(z))
266
+ return -ops.mean(ops.log(real + EPS))
267
+
268
+ def discriminator_loss(self, z):
269
+ r"""Computes the loss of the discriminator."""
270
+ real = ops.sigmoid(self.discriminator(_randn_like(z)))
271
+ fake = ops.sigmoid(self.discriminator(ops.stop_gradient(z)))
272
+ real_loss = -ops.mean(ops.log(real + EPS))
273
+ fake_loss = -ops.mean(ops.log(1 - fake + EPS))
274
+ return real_loss + fake_loss
275
+
276
+
277
+ class ARGVA(ARGA):
278
+ r"""The Adversarially Regularized Variational Graph Auto-Encoder model
279
+ from the `"Adversarially Regularized Graph Autoencoder for Graph
280
+ Embedding" <https://arxiv.org/abs/1802.04407>`_ paper.
281
+
282
+ Args:
283
+ encoder: The encoder module to compute :math:`\mu` and
284
+ :math:`\log\sigma^2`.
285
+ discriminator: The discriminator module.
286
+ decoder (optional): The decoder module. If set to `None`, will
287
+ default to `InnerProductDecoder`. (default: `None`)
288
+ """
289
+ def __init__(self, encoder, discriminator, decoder=None):
290
+ super().__init__(encoder, discriminator, decoder)
291
+ self.vgae = VGAE(encoder, decoder)
292
+
293
+ @property
294
+ def _mu(self):
295
+ return self.vgae._mu
296
+
297
+ @property
298
+ def _logstd(self):
299
+ return self.vgae._logstd
300
+
301
+ def reparametrize(self, mu, logstd):
302
+ return self.vgae.reparametrize(mu, logstd)
303
+
304
+ def encode(self, *args, **kwargs):
305
+ return self.vgae.encode(*args, **kwargs)
306
+
307
+ def kl_loss(self, mu=None, logstd=None):
308
+ return self.vgae.kl_loss(mu, logstd)
309
+
310
+ def eval(self):
311
+ self.training = False
312
+ self.vgae.eval()
313
+ return self
314
+
315
+ def train(self):
316
+ self.training = True
317
+ self.vgae.train()
318
+ return self