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,182 @@
1
+ import keras
2
+ from keras import ops
3
+ from k3_node.layers.conv.message_passing import MessagePassing
4
+ from k3_node.layers.conv.utils import scatter, softmax
5
+ from k3_node.layers.conv.utils import is_tracing
6
+
7
+
8
+ class HypergraphConv(MessagePassing):
9
+ r"""The hypergraph convolutional operator from the `"Hypergraph Convolution
10
+ and Hypergraph Attention" <https://arxiv.org/abs/1901.08150>`_ paper.
11
+
12
+ Args:
13
+ in_channels (int): Size of each input sample.
14
+ out_channels (int): Size of each output sample.
15
+ use_attention (bool, optional): Whether to use hypergraph attention. (default: :obj:`False`)
16
+ attention_mode (str, optional): Attention mode (:obj:`"node"` or :obj:`"edge"`). (default: :obj:`"node"`)
17
+ heads (int, optional): Number of multi-head-attentions. (default: :obj:`1`)
18
+ concat (bool, optional): Whether to concatenate heads. (default: :obj:`True`)
19
+ negative_slope (float, optional): LeakyReLU angle. (default: :obj:`0.2`)
20
+ bias (bool, optional): Whether to learn an additive bias. (default: :obj:`True`)
21
+
22
+ Example:
23
+ ```python
24
+ import numpy as np
25
+ from k3_node.layers import HypergraphConv
26
+
27
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
28
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
29
+
30
+ # Row 0: node index, row 1: hyperedge index (nodes 0-5 form two hyperedges)
31
+ hyperedge_index = np.array([[0, 1, 2, 3, 4, 5], [0, 0, 0, 1, 1, 1]])
32
+ layer = HypergraphConv(in_channels=8, out_channels=16)
33
+ out = layer(x, hyperedge_index)
34
+ print(tuple(out.shape)) # (10, 16)
35
+ ```
36
+ """
37
+
38
+ def __init__(
39
+ self,
40
+ in_channels: int,
41
+ out_channels: int,
42
+ use_attention: bool = False,
43
+ attention_mode: str = "node",
44
+ heads: int = 1,
45
+ concat: bool = True,
46
+ negative_slope: float = 0.2,
47
+ bias: bool = True,
48
+ **kwargs,
49
+ ):
50
+ kwargs.setdefault("aggr", "add")
51
+ super().__init__(node_dim=0, **kwargs)
52
+
53
+ assert attention_mode in ["node", "edge"]
54
+
55
+ self.in_channels = in_channels
56
+ self.out_channels = out_channels
57
+ self.use_attention = use_attention
58
+ self.attention_mode = attention_mode
59
+ self.negative_slope = negative_slope
60
+ self.use_bias = bias
61
+
62
+ if self.use_attention:
63
+ self.heads = heads
64
+ self.concat = concat
65
+ self.lin = keras.layers.Dense(heads * out_channels, use_bias=False)
66
+ else:
67
+ self.heads = 1
68
+ self.concat = True
69
+ self.lin = keras.layers.Dense(out_channels, use_bias=False)
70
+
71
+ def build(self, input_shape=None):
72
+ if not self.lin.built and self.in_channels > 0:
73
+ self.lin.build((None, self.in_channels))
74
+
75
+ if self.use_attention:
76
+ self.att = self.add_weight(
77
+ shape=(1, self.heads, 2 * self.out_channels),
78
+ initializer="glorot_uniform",
79
+ trainable=True,
80
+ name="att",
81
+ )
82
+
83
+ out_dim = self.heads * self.out_channels if self.concat else self.out_channels
84
+ if self.use_bias:
85
+ self.bias = self.add_weight(
86
+ shape=(out_dim,),
87
+ initializer="zeros",
88
+ trainable=True,
89
+ name="bias",
90
+ )
91
+ else:
92
+ self.bias = None
93
+
94
+ super().build(input_shape)
95
+
96
+ def call(
97
+ self,
98
+ x,
99
+ hyperedge_index,
100
+ hyperedge_weight=None,
101
+ hyperedge_attr=None,
102
+ num_edges=None,
103
+ ):
104
+ if not self.built:
105
+ self.build((None, self.in_channels))
106
+
107
+ num_nodes = ops.shape(x)[0]
108
+ if num_edges is None:
109
+ if not is_tracing(hyperedge_index):
110
+ try:
111
+ num_edges = int(ops.max(hyperedge_index[1])) + 1
112
+ except Exception:
113
+ num_edges = None
114
+
115
+ if hyperedge_weight is None:
116
+ edge_w = ops.ones((ops.shape(hyperedge_index)[1],), dtype=x.dtype)
117
+ else:
118
+ edge_w = ops.take(hyperedge_weight, hyperedge_index[1], axis=0)
119
+
120
+ x = self.lin(x)
121
+
122
+ alpha = None
123
+ if self.use_attention:
124
+ assert hyperedge_attr is not None
125
+ x = ops.reshape(x, (-1, self.heads, self.out_channels))
126
+ hyperedge_attr = self.lin(hyperedge_attr)
127
+ hyperedge_attr = ops.reshape(hyperedge_attr, (-1, self.heads, self.out_channels))
128
+
129
+ x_i = ops.take(x, hyperedge_index[0], axis=0)
130
+ x_j = ops.take(hyperedge_attr, hyperedge_index[1], axis=0)
131
+
132
+ alpha = ops.sum(ops.concatenate([x_i, x_j], axis=-1) * self.att, axis=-1)
133
+ alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
134
+
135
+ if self.attention_mode == "node":
136
+ alpha = softmax(alpha, index=hyperedge_index[1], num_nodes=num_edges)
137
+ else:
138
+ alpha = softmax(alpha, index=hyperedge_index[0], num_nodes=num_nodes)
139
+
140
+ D = scatter(edge_w, hyperedge_index[0], dim=0, dim_size=num_nodes, reduce="sum")
141
+ D = ops.where(ops.equal(D, 0), 0.0, 1.0 / D)
142
+
143
+ ones_e = ops.ones((ops.shape(hyperedge_index)[1],), dtype=x.dtype)
144
+ B = scatter(ones_e, hyperedge_index[1], dim=0, dim_size=num_edges, reduce="sum")
145
+ B = ops.where(ops.equal(B, 0), 0.0, 1.0 / B)
146
+
147
+ # 1. nodes -> hyperedges
148
+ out = self.propagate(
149
+ hyperedge_index,
150
+ x=x,
151
+ norm=B,
152
+ alpha=alpha,
153
+ size=(num_nodes, num_edges),
154
+ )
155
+ # 2. hyperedges -> nodes
156
+ flipped_edge_index = ops.stack([hyperedge_index[1], hyperedge_index[0]], axis=0)
157
+ out = self.propagate(
158
+ flipped_edge_index,
159
+ x=out,
160
+ norm=D,
161
+ alpha=alpha,
162
+ size=(num_edges, num_nodes),
163
+ )
164
+
165
+ if self.concat:
166
+ out = ops.reshape(out, (-1, self.heads * self.out_channels))
167
+ else:
168
+ out = ops.mean(out, axis=1)
169
+
170
+ if self.bias is not None:
171
+ out = out + self.bias
172
+
173
+ return out
174
+
175
+ def message(self, x_j, norm_i, alpha=None):
176
+ H, F = self.heads, self.out_channels
177
+ norm_i_exp = ops.expand_dims(ops.expand_dims(norm_i, axis=-1), axis=-1)
178
+ out = norm_i_exp * ops.reshape(x_j, (-1, H, F))
179
+ if alpha is not None:
180
+ out = ops.expand_dims(alpha, axis=-1) * out
181
+ return out
182
+
@@ -0,0 +1,81 @@
1
+ from typing import Union, Tuple
2
+ from keras import layers, ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class LEConv(MessagePassing):
8
+ r"""The local extremum graph convolutional operator from the
9
+ `"ASAP: Adaptive Structure Aware Pooling for Learning Hierarchical Graph
10
+ Representations" <https://arxiv.org/abs/1911.07979>`_ paper.
11
+
12
+ Args:
13
+ in_channels: Size of each input sample, or a tuple for bipartite graphs.
14
+ out_channels: Size of each output sample.
15
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
16
+ (default: ``True``)
17
+
18
+ Example:
19
+ ```python
20
+ import numpy as np
21
+ from k3_node.layers import LEConv
22
+
23
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
24
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
25
+
26
+ layer = LEConv(in_channels=8, out_channels=16)
27
+ out = layer(x, edge_index)
28
+ print(tuple(out.shape)) # (10, 16)
29
+ ```
30
+ """
31
+
32
+ def __init__(
33
+ self,
34
+ in_channels: Union[int, Tuple[int, int]],
35
+ out_channels: int,
36
+ bias: bool = True,
37
+ **kwargs,
38
+ ):
39
+ super().__init__(aggr="add", **kwargs)
40
+ self.in_channels = in_channels
41
+ self.out_channels = out_channels
42
+ self.use_bias = bias
43
+
44
+ self.lin1 = layers.Dense(out_channels, use_bias=bias)
45
+ self.lin2 = layers.Dense(out_channels, use_bias=False)
46
+ self.lin3 = layers.Dense(out_channels, use_bias=bias)
47
+
48
+ def build(self, input_shape):
49
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
50
+ in_channels_src = input_shape[0][-1]
51
+ in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
52
+ else:
53
+ in_channels_src = input_shape[-1]
54
+ in_channels_dst = input_shape[-1]
55
+
56
+ self.lin1.build((None, in_channels_src))
57
+ self.lin2.build((None, in_channels_dst))
58
+ self.lin3.build((None, in_channels_dst))
59
+ self.built = True
60
+
61
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
62
+ if edge_index is None and isinstance(x, (tuple, list)):
63
+ x, edge_index = x[0], x[1]
64
+
65
+ if not isinstance(x, (tuple, list)):
66
+ x_src, x_dst = x, x
67
+ else:
68
+ x_src, x_dst = x[0], x[1]
69
+
70
+ a = self.lin1(x_src)
71
+ b = self.lin2(x_dst)
72
+
73
+ out = self.propagate(edge_index, a=a, b=b, edge_weight=edge_weight)
74
+ return out + self.lin3(x_dst)
75
+
76
+ def message(self, a_j, b_i, edge_weight=None):
77
+ out = a_j - b_i
78
+ if edge_weight is None:
79
+ return out
80
+ return out * ops.expand_dims(edge_weight, -1)
81
+
@@ -0,0 +1,58 @@
1
+ from keras import ops
2
+
3
+ from k3_node.layers.conv.message_passing import MessagePassing
4
+ from k3_node.layers.conv.utils import gcn_norm
5
+
6
+
7
+ class LGConv(MessagePassing):
8
+ r"""The LightGCN operator from the `"LightGCN: Simplifying and Powering
9
+ Graph Convolution Network for Recommendation"
10
+ <https://arxiv.org/abs/2002.02126>`_ paper.
11
+
12
+ Args:
13
+ normalize: Whether to apply symmetric normalization. (default: ``True``)
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import LGConv
19
+
20
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
21
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
22
+
23
+ layer = LGConv()
24
+ out = layer(x, edge_index)
25
+ print(tuple(out.shape)) # (10, 8)
26
+ ```
27
+ """
28
+
29
+ weighted_sum_message = True
30
+
31
+ def __init__(self, normalize: bool = True, **kwargs):
32
+ super().__init__(aggr="add", **kwargs)
33
+ self.normalize = normalize
34
+
35
+ def build(self, input_shape):
36
+ self.built = True
37
+
38
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
39
+ if edge_index is None and isinstance(x, (tuple, list)):
40
+ x, edge_index = x[0], x[1]
41
+
42
+ if self.normalize:
43
+ num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
44
+ edge_index, edge_weight = gcn_norm(
45
+ edge_index,
46
+ edge_weight,
47
+ num_nodes=num_nodes,
48
+ add_self_loops=False,
49
+ flow=self.flow,
50
+ dtype=x.dtype,
51
+ )
52
+
53
+ return self.propagate(edge_index, x=x, edge_weight=edge_weight)
54
+
55
+ def message(self, x_j, edge_weight=None):
56
+ if edge_weight is None:
57
+ return x_j
58
+ return ops.expand_dims(edge_weight, -1) * x_j
@@ -0,0 +1,84 @@
1
+ from typing import Optional, List
2
+ import keras
3
+ from keras import ops
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class MeshCNNConv(MessagePassing):
8
+ r"""The MeshCNN convolutional operator from the `"MeshCNN: A Network With An Edge"
9
+ <https://arxiv.org/abs/1809.05910>`_ paper.
10
+
11
+ Args:
12
+ in_channels (int): Size of each input sample.
13
+ out_channels (int): Size of each output sample.
14
+ kernels (List[keras.layers.Layer], optional): A list of 5 neural network layers
15
+ that transform edge representations. (default: :obj:`None`)
16
+
17
+ Example:
18
+ ```python
19
+ import numpy as np
20
+ from k3_node.layers import MeshCNNConv
21
+
22
+ # In MeshCNN the "nodes" are mesh edges; each mesh edge has exactly 4 neighboring edges,
23
+ # listed in a fixed order in edge_index (4 incoming entries per mesh edge).
24
+ x = np.random.rand(4, 8).astype("float32") # features of 4 mesh edges
25
+ edge_index = np.array([
26
+ [1, 2, 3, 0, 0, 2, 3, 1, 0, 1, 3, 2, 0, 1, 2, 3], # neighboring mesh edge
27
+ [0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3], # mesh edge being updated
28
+ ])
29
+ layer = MeshCNNConv(in_channels=8, out_channels=16)
30
+ out = layer(x, edge_index)
31
+ print(tuple(out.shape)) # (4, 16)
32
+ ```
33
+ """
34
+
35
+ def __init__(
36
+ self,
37
+ in_channels: int,
38
+ out_channels: int,
39
+ kernels: Optional[List[keras.layers.Layer]] = None,
40
+ **kwargs,
41
+ ):
42
+ kwargs.setdefault("aggr", "add")
43
+ super().__init__(**kwargs)
44
+
45
+ self.in_channels = in_channels
46
+ self.out_channels = out_channels
47
+
48
+ if kernels is None:
49
+ self.kernels = [
50
+ keras.layers.Dense(out_channels, use_bias=True)
51
+ for _ in range(5)
52
+ ]
53
+ else:
54
+ assert len(kernels) == 5, "kernels must be a list of 5 layers"
55
+ self.kernels = kernels
56
+
57
+ def build(self, input_shape=None):
58
+ for k in self.kernels:
59
+ if hasattr(k, "build") and not k.built:
60
+ k.build((None, self.in_channels))
61
+ super().build(input_shape)
62
+
63
+ def call(self, x, edge_index, **kwargs):
64
+ return self.propagate(edge_index, x=x)
65
+
66
+ def message(self, x_j):
67
+ n_a = x_j[0::4]
68
+ n_b = x_j[1::4]
69
+ n_c = x_j[2::4]
70
+ n_d = x_j[3::4]
71
+
72
+ m1 = self.kernels[1](ops.abs(n_a - n_c))
73
+ m2 = self.kernels[2](n_a + n_c)
74
+ m3 = self.kernels[3](ops.abs(n_b - n_d))
75
+ m4 = self.kernels[4](n_b + n_d)
76
+
77
+ return ops.reshape(
78
+ ops.stack([m1, m2, m3, m4], axis=1),
79
+ (-1, self.out_channels),
80
+ )
81
+
82
+ def update(self, inputs, x):
83
+ return self.kernels[0](x) + inputs
84
+