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,403 @@
1
+ import keras
2
+ from typing import Optional, Tuple
3
+ from keras import layers, ops
4
+ import numpy as np
5
+ from k3_node.ops.segment import segment_max, segment_sum
6
+ from k3_node.ops.creation import scatter, full
7
+
8
+
9
+ def ptr2index(ptr):
10
+ r"""Converts a pointer tensor into an index tensor.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import ptr2index
16
+
17
+ ptr = np.array([0, 3, 5]) # CSR pointer: set 0 has 3 elements, set 1 has 2
18
+ print(tuple(ptr2index(ptr).shape)) # (5,): set id of every element
19
+ ```
20
+ """
21
+ ptr_np = ops.convert_to_numpy(ptr).astype(np.int64)
22
+ counts = ptr_np[1:] - ptr_np[:-1]
23
+ index_np = np.repeat(np.arange(len(counts), dtype=np.int64), counts)
24
+ return ops.convert_to_tensor(index_np, dtype="int32")
25
+
26
+
27
+ def to_dense_batch(
28
+ x,
29
+ index,
30
+ dim_size: Optional[int] = None,
31
+ fill_value: float = 0.0,
32
+ max_num_elements: Optional[int] = None,
33
+ ) -> Tuple[any, any]:
34
+ r"""Transforms a batched feature tensor into a dense representation
35
+ of shape `(batch_size, max_nodes, *dims)`.
36
+
37
+ Example:
38
+ ```python
39
+ import numpy as np
40
+ from k3_node.layers import to_dense_batch
41
+
42
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
43
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
44
+
45
+ x_dense, mask = to_dense_batch(x, batch) # [num_graphs, max_nodes, features] + validity mask
46
+ print(tuple(x_dense.shape), tuple(mask.shape)) # (2, 5, 8) (2, 5)
47
+ ```
48
+ """
49
+ from k3_node.layers.conv.utils import is_tracing
50
+ from k3_node.ops.host import _in_shape_inference
51
+
52
+ # Keras' shape inference traces too, but only needs shapes: the host path then sees zeros.
53
+ if not is_tracing(index) or _in_shape_inference():
54
+ try:
55
+ # `index` is purely structural (never differentiated), so plain
56
+ # numpy is fine for it. `x` itself is placed into the dense
57
+ # tensor with the differentiable `ops.scatter` below -- a numpy
58
+ # round-trip on `x` would silently detach it from the graph and
59
+ # stop gradients from flowing back into whatever produced it.
60
+ from k3_node.ops.host import to_numpy # zeros during Keras' shape inference
61
+
62
+ index_np = np.asarray(to_numpy(index)).astype(np.int64)
63
+ N = len(index_np)
64
+
65
+ B = int(np.max(index_np)) + 1 if N > 0 else 0
66
+ if dim_size is not None:
67
+ B = max(B, int(dim_size))
68
+
69
+ # Compute local index for each node in its graph. `index` is
70
+ # required to be sorted (see `assert_sorted_index`), so the
71
+ # first occurrence of each value is its graph's start offset.
72
+ counts = np.bincount(index_np, minlength=B)
73
+ max_nodes = int(np.max(counts)) if len(counts) > 0 else 0
74
+ if max_num_elements is not None:
75
+ max_nodes = max(max_nodes, int(max_num_elements))
76
+
77
+ local_index = np.arange(N) - np.searchsorted(index_np, index_np, side="left")
78
+ valid = local_index < max_nodes
79
+
80
+ mask_np = np.zeros((B, max_nodes), dtype=bool)
81
+ mask_np[index_np[valid], local_index[valid]] = True
82
+
83
+ feat_shape = tuple(ops.shape(x)[1:])
84
+ scatter_idx = np.stack([index_np[valid], local_index[valid]], axis=1)
85
+ x_valid = x if bool(valid.all()) else ops.take(x, np.nonzero(valid)[0], axis=0)
86
+ out = scatter(scatter_idx, x_valid, shape=(B, max_nodes, *feat_shape))
87
+
88
+ if fill_value != 0.0:
89
+ mask_t = ops.convert_to_tensor(mask_np)
90
+ mask_expanded = ops.reshape(mask_t, (B, max_nodes) + (1,) * len(feat_shape))
91
+ fill = full((B, max_nodes, *feat_shape), fill_value, dtype=x.dtype)
92
+ out = ops.where(mask_expanded, out, fill)
93
+
94
+ return out, ops.convert_to_tensor(mask_np, dtype="bool")
95
+ except Exception:
96
+ pass
97
+
98
+ # Pure ops implementation for symbolic tracing / graph execution. `index` is sorted, so a
99
+ # node's position in its graph is its offset from the graph's first node (O(N) memory).
100
+ from k3_node.ops.segment import segment_sum
101
+
102
+ N = ops.shape(x)[0]
103
+ index = ops.cast(index, "int32")
104
+ static_size = isinstance(dim_size, (int, np.integer)) # a tensor while tracing
105
+ num_segments = int(dim_size) if static_size else N
106
+ counts = segment_sum(ops.ones_like(index), index, num_segments=num_segments)
107
+ starts = ops.cumsum(counts) - counts
108
+ local_idx = ops.arange(N, dtype="int32") - ops.take(starts, index, axis=0)
109
+
110
+ if max_num_elements is not None:
111
+ max_nodes = int(max_num_elements)
112
+ elif hasattr(x, "shape") and x.shape[0] is not None:
113
+ max_nodes = int(x.shape[0])
114
+ else:
115
+ max_nodes = ops.max(local_idx) + 1
116
+
117
+ B = int(dim_size) if static_size else ops.max(index) + 1
118
+ from k3_node.layers.conv.utils import is_tracing
119
+
120
+ if not static_size and keras.config.backend() == "jax" and is_tracing(index):
121
+ raise ValueError(
122
+ "Compiled JAX code needs the number of graphs as a Python int to build the dense batch. "
123
+ "Pass it explicitly, e.g. `batch_size=data.num_graphs` (or `dim_size=` for "
124
+ "`to_dense_batch`), or train with `run_eagerly=True`."
125
+ )
126
+
127
+ dense_x = full((B, max_nodes, *ops.shape(x)[1:]), fill_value, dtype=x.dtype)
128
+ mask = ops.zeros((B, max_nodes), dtype="bool")
129
+
130
+ scatter_indices = ops.stack([ops.cast(index, "int32"), ops.cast(local_idx, "int32")], axis=1)
131
+ dense_x = ops.scatter_update(dense_x, scatter_indices, x)
132
+ mask = ops.scatter_update(mask, scatter_indices, ops.ones((N,), dtype="bool"))
133
+ return dense_x, mask
134
+
135
+
136
+ def from_dense_batch(x_dense, index):
137
+ r"""The inverse of :func:`to_dense_batch`: turns ``[batch_size, max_nodes, *dims]`` back into
138
+ one row per node, in the order given by the sorted ``index``. Static shapes, differentiable.
139
+
140
+ Example:
141
+ ```python
142
+ import numpy as np
143
+ from k3_node.layers import from_dense_batch, to_dense_batch
144
+
145
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
146
+ batch = np.array([0, 0, 0, 1, 1, 1, 1, 1, 2, 2]) # three graphs of different sizes
147
+
148
+ x_dense, mask = to_dense_batch(x, batch)
149
+ print(np.allclose(from_dense_batch(x_dense, batch), x)) # True
150
+ ```
151
+ """
152
+ from k3_node.ops.segment import segment_sum
153
+
154
+ index = ops.cast(ops.convert_to_tensor(index), "int32")
155
+ num_graphs, max_nodes = ops.shape(x_dense)[0], ops.shape(x_dense)[1]
156
+ counts = segment_sum(ops.ones_like(index), index, num_segments=num_graphs)
157
+ starts = ops.cumsum(counts) - counts
158
+ position = ops.arange(ops.shape(index)[0], dtype="int32") - ops.take(starts, index, axis=0)
159
+ flat = ops.reshape(x_dense, (-1,) + tuple(x_dense.shape[2:]))
160
+ return ops.take(flat, index * max_nodes + position, axis=0)
161
+
162
+
163
+ def to_dense_adj(edge_index, batch=None, edge_attr=None, max_num_nodes: Optional[int] = None,
164
+ batch_size: Optional[int] = None):
165
+ r"""Converts a batch of graphs into dense adjacency matrices of shape
166
+ ``[num_graphs, max_nodes, max_nodes]`` (or ``[..., edge_features]`` with ``edge_attr``), as in PyG.
167
+
168
+ The node dimension matches :func:`to_dense_batch` for the same ``batch`` vector, so the two can
169
+ be used together. Duplicate edges are summed.
170
+
171
+ Example:
172
+ ```python
173
+ import numpy as np
174
+ from k3_node.layers import to_dense_adj
175
+
176
+ edge_index = np.array([[0, 1, 2, 3, 4], [1, 2, 0, 4, 3]])
177
+ batch = np.array([0, 0, 0, 1, 1]) # two graphs with 3 and 2 nodes
178
+ adj = to_dense_adj(edge_index, batch)
179
+ print(tuple(adj.shape)) # (2, 3, 3)
180
+ ```
181
+ """
182
+ from k3_node.layers.conv.utils import is_tracing
183
+ from k3_node.ops.host import _in_shape_inference
184
+
185
+ edge_index = ops.cast(ops.convert_to_tensor(edge_index), "int32")
186
+ if batch is None:
187
+ num_nodes = max_num_nodes or (int(ops.convert_to_numpy(ops.max(edge_index))) + 1)
188
+ batch = ops.zeros((num_nodes,), dtype="int32")
189
+ batch = ops.cast(ops.convert_to_tensor(batch), "int32")
190
+
191
+ if not is_tracing(batch) or (_in_shape_inference() and None not in tuple(batch.shape)):
192
+ from k3_node.ops.host import to_numpy # zeros during Keras' shape inference
193
+
194
+ index_np = np.asarray(to_numpy(batch)).astype(np.int64)
195
+ num_graphs = int(index_np.max()) + 1 if len(index_np) else 0
196
+ num_graphs = max(num_graphs, int(batch_size or 0))
197
+ local = np.arange(len(index_np)) - np.searchsorted(index_np, index_np, side="left")
198
+ max_nodes = int(np.bincount(index_np, minlength=num_graphs).max()) if len(index_np) else 0
199
+ max_nodes = max(max_nodes, int(max_num_nodes or 0))
200
+ local = ops.convert_to_tensor(local.astype("int32"))
201
+ else: # compiled: the same local indices and padding as to_dense_batch's compiled path
202
+ from k3_node.ops.segment import segment_sum
203
+
204
+ n = ops.shape(batch)[0]
205
+ static = isinstance(batch_size, (int, np.integer))
206
+ counts = segment_sum(ops.ones_like(batch), batch, num_segments=int(batch_size) if static else n)
207
+ local = ops.arange(n, dtype="int32") - ops.take(ops.cumsum(counts) - counts, batch, axis=0)
208
+ num_graphs = batch_size if static else ops.max(batch) + 1
209
+ if max_num_nodes is not None:
210
+ max_nodes = int(max_num_nodes)
211
+ elif batch.shape[0] is not None:
212
+ max_nodes = int(batch.shape[0])
213
+ else:
214
+ max_nodes = ops.max(local) + 1
215
+
216
+ src, dst = edge_index[0], edge_index[1]
217
+ indices = ops.stack([ops.take(batch, src), ops.take(local, src), ops.take(local, dst)], axis=1)
218
+ num_edges = ops.shape(edge_index)[1]
219
+ values = ops.ones((num_edges,), dtype="float32") if edge_attr is None else ops.convert_to_tensor(edge_attr)
220
+ extra = tuple(values.shape[1:])
221
+ return scatter(indices, values, (num_graphs, max_nodes, max_nodes) + extra)
222
+
223
+
224
+ class Aggregation(layers.Layer):
225
+ r"""An abstract base class for implementing custom aggregations."""
226
+
227
+ def __init__(self, **kwargs):
228
+ super().__init__(**kwargs)
229
+ self.built = True
230
+
231
+ def build(self, input_shape=None):
232
+ self.built = True
233
+
234
+ def reset_parameters(self):
235
+ r"""Resets all learnable parameters of the module."""
236
+ pass
237
+
238
+ def __call__(
239
+ self,
240
+ x,
241
+ index: Optional[any] = None,
242
+ ptr: Optional[any] = None,
243
+ dim_size: Optional[int] = None,
244
+ dim: int = -2,
245
+ max_num_elements: Optional[int] = None,
246
+ **kwargs,
247
+ ):
248
+ # Plain NumPy inputs cannot be mixed with backend tensors (e.g. `ndarray - torch.Tensor`).
249
+ if isinstance(x, np.ndarray):
250
+ x = ops.convert_to_tensor(x)
251
+ dim_total = len(x.shape) if hasattr(x, "shape") and x.shape is not None else len(ops.shape(x))
252
+ if dim >= dim_total or dim < -dim_total:
253
+ raise ValueError(
254
+ f"Encountered invalid dimension '{dim}' of source tensor with "
255
+ f"{dim_total} dimensions"
256
+ )
257
+
258
+ if index is None and ptr is None:
259
+ N = x.shape[dim] if hasattr(x, "shape") and x.shape[dim] is not None else ops.shape(x)[dim]
260
+ index = ops.zeros((N,), dtype="int32")
261
+
262
+ if ptr is not None and index is None:
263
+ index = ptr2index(ptr)
264
+
265
+ if ptr is not None:
266
+ ptr_len = ptr.shape[0] if hasattr(ptr, "shape") and ptr.shape[0] is not None else ops.shape(ptr)[0]
267
+ if dim_size is None:
268
+ dim_size = ptr_len - 1
269
+ elif dim_size != ptr_len - 1:
270
+ raise ValueError(
271
+ f"Encountered invalid 'dim_size' (got '{dim_size}' but "
272
+ f"expected '{ptr_len - 1}')"
273
+ )
274
+
275
+ if index is not None and dim_size is None:
276
+ dim_size = ops.max(index) + 1
277
+ try:
278
+ dim_size = int(dim_size)
279
+ except Exception:
280
+ pass
281
+
282
+ # Handle positional / keyword call to call()
283
+ return self.call(
284
+ x,
285
+ index=index,
286
+ ptr=ptr,
287
+ dim_size=dim_size,
288
+ dim=dim,
289
+ max_num_elements=max_num_elements,
290
+ **kwargs,
291
+ )
292
+
293
+ def call(
294
+ self,
295
+ x,
296
+ index: Optional[any] = None,
297
+ ptr: Optional[any] = None,
298
+ dim_size: Optional[int] = None,
299
+ dim: int = -2,
300
+ max_num_elements: Optional[int] = None,
301
+ ):
302
+ raise NotImplementedError
303
+
304
+ def assert_index_present(self, index: Optional[any]):
305
+ if index is None:
306
+ raise NotImplementedError("Aggregation requires 'index' to be specified")
307
+
308
+ def assert_sorted_index(self, index: Optional[any]):
309
+ if index is not None:
310
+ from k3_node.layers.conv.utils import is_tracing
311
+ if is_tracing(index):
312
+ return
313
+ idx_np = ops.convert_to_numpy(index)
314
+ if not np.all(idx_np[:-1] <= idx_np[1:]):
315
+ raise ValueError(
316
+ "Can not perform aggregation since the 'index' tensor is not sorted. "
317
+ "Specifically, if you use this aggregation as part of 'MessagePassing', "
318
+ "ensure that 'edge_index' is sorted by destination nodes."
319
+ )
320
+
321
+ def assert_two_dimensional_input(self, x, dim: int = -2):
322
+ if len(ops.shape(x)) != 2:
323
+ raise ValueError(
324
+ f"Aggregation requires two-dimensional inputs (got '{len(ops.shape(x))}')"
325
+ )
326
+ if dim not in [-2, 0]:
327
+ raise ValueError(
328
+ f"Aggregation needs to perform aggregation in first dimension (got '{dim}')"
329
+ )
330
+
331
+ def reduce(
332
+ self,
333
+ x,
334
+ index: Optional[any] = None,
335
+ ptr: Optional[any] = None,
336
+ dim_size: Optional[int] = None,
337
+ dim: int = -2,
338
+ reduce: str = "sum",
339
+ ):
340
+ r"""Reduces features along groups specified by `index` or `ptr`."""
341
+ if ptr is not None and index is None:
342
+ index = ptr2index(ptr)
343
+
344
+ if index is None:
345
+ raise RuntimeError("Aggregation requires 'index' to be specified")
346
+
347
+ index = ops.cast(index, dtype="int32")
348
+ if dim_size is None:
349
+ dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
350
+
351
+ if reduce in ["sum", "add"]:
352
+ return segment_sum(x, index, num_segments=dim_size)
353
+ elif reduce == "mean":
354
+ sum_val = segment_sum(x, index, num_segments=dim_size)
355
+ ones = ops.ones_like(x)
356
+ count = segment_sum(ones, index, num_segments=dim_size)
357
+ return sum_val / ops.maximum(count, 1.0)
358
+ elif reduce == "max":
359
+ val = segment_max(x, index, num_segments=dim_size)
360
+ ones = ops.ones_like(x)
361
+ count = segment_sum(ones, index, num_segments=dim_size)
362
+ return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
363
+ elif reduce == "min":
364
+ val = -segment_max(-x, index, num_segments=dim_size)
365
+ ones = ops.ones_like(x)
366
+ count = segment_sum(ones, index, num_segments=dim_size)
367
+ return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
368
+ elif reduce == "mul":
369
+ log_abs = ops.log(ops.maximum(ops.abs(x), 1e-7))
370
+ sum_log = segment_sum(log_abs, index, num_segments=dim_size)
371
+ neg_count = segment_sum(ops.cast(ops.less(x, 0.0), dtype=x.dtype), index, num_segments=dim_size)
372
+ sign = ops.cos(ops.cast(3.141592653589793, dtype=x.dtype) * neg_count)
373
+ return ops.exp(sum_log) * sign
374
+ else:
375
+ raise ValueError(f"Unsupported reduction '{reduce}'")
376
+
377
+ def to_dense_batch(
378
+ self,
379
+ x,
380
+ index: Optional[any] = None,
381
+ ptr: Optional[any] = None,
382
+ dim_size: Optional[int] = None,
383
+ dim: int = -2,
384
+ fill_value: float = 0.0,
385
+ max_num_elements: Optional[int] = None,
386
+ ) -> Tuple[any, any]:
387
+ if ptr is not None and index is None:
388
+ index = ptr2index(ptr)
389
+
390
+ self.assert_index_present(index)
391
+ self.assert_sorted_index(index)
392
+ self.assert_two_dimensional_input(x, dim)
393
+
394
+ return to_dense_batch(
395
+ x,
396
+ index,
397
+ dim_size=dim_size,
398
+ fill_value=fill_value,
399
+ max_num_elements=max_num_elements,
400
+ )
401
+
402
+ def __repr__(self) -> str:
403
+ return f"{self.__class__.__name__}()"