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,327 @@
1
+ import math
2
+ from typing import Any, Dict, List, Optional, Union
3
+ from keras import initializers, layers, ops
4
+
5
+
6
+ def _get_weight_initializer(initializer: Optional[str], in_channels: int):
7
+ if initializer in ('glorot', 'xavier_uniform'):
8
+ return initializers.GlorotUniform()
9
+ elif initializer in ('kaiming_uniform', 'he_uniform'):
10
+ return initializers.HeUniform()
11
+ elif initializer == 'uniform' or initializer is None:
12
+ if in_channels > 0:
13
+ bound = 1.0 / math.sqrt(in_channels)
14
+ return initializers.RandomUniform(minval=-bound, maxval=bound)
15
+ return initializers.GlorotUniform()
16
+ return initializers.get(initializer)
17
+
18
+
19
+ def _get_bias_initializer(initializer: Optional[str], in_channels: int):
20
+ if initializer == 'zeros':
21
+ return initializers.Zeros()
22
+ elif initializer is None:
23
+ if in_channels > 0:
24
+ bound = 1.0 / math.sqrt(in_channels)
25
+ return initializers.RandomUniform(minval=-bound, maxval=bound)
26
+ return initializers.Zeros()
27
+ return initializers.get(initializer)
28
+
29
+
30
+ class Linear(layers.Layer):
31
+ r"""Applies a linear transformation to the incoming data:
32
+
33
+ .. math::
34
+ \mathbf{x}^{\prime} = \mathbf{x} \mathbf{W} + \mathbf{b}
35
+
36
+ Args:
37
+ in_channels (int): Size of each input sample.
38
+ out_channels (int): Size of each output sample.
39
+ bias (bool, optional): If set to :obj:`False`, the layer will not learn
40
+ an additive bias. (default: :obj:`True`)
41
+ weight_initializer (str, optional): The initializer for the weight
42
+ matrix (:obj:`"glorot"`, :obj:`"uniform"`, :obj:`"kaiming_uniform"`
43
+ or :obj:`None`). (default: :obj:`None`)
44
+ bias_initializer (str, optional): The initializer for the bias vector
45
+ (:obj:`"zeros"` or :obj:`None`). (default: :obj:`None`)
46
+
47
+ Example:
48
+ ```python
49
+ import numpy as np
50
+ from k3_node.layers import Linear
51
+
52
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
53
+
54
+ layer = Linear(in_channels=8, out_channels=16)
55
+ print(tuple(layer(x).shape)) # (10, 16)
56
+ ```
57
+ """
58
+ def __init__(
59
+ self,
60
+ in_channels: int,
61
+ out_channels: int,
62
+ bias: bool = True,
63
+ weight_initializer: Optional[str] = None,
64
+ bias_initializer: Optional[str] = None,
65
+ **kwargs
66
+ ):
67
+ super().__init__(**kwargs)
68
+ self.in_channels = in_channels
69
+ self.out_channels = out_channels
70
+ self.use_bias = bias
71
+ self.weight_initializer = weight_initializer
72
+ self.bias_initializer = bias_initializer
73
+ self.weight = None
74
+ self.bias = None
75
+
76
+ if self.in_channels > 0:
77
+ self._build_weights(self.in_channels)
78
+
79
+ def _build_weights(self, in_features: int):
80
+ self.in_channels = in_features
81
+ w_init = _get_weight_initializer(self.weight_initializer, in_features)
82
+ b_init = _get_bias_initializer(self.bias_initializer, in_features)
83
+
84
+ self.weight = self.add_weight(
85
+ shape=(in_features, self.out_channels),
86
+ initializer=w_init,
87
+ trainable=True,
88
+ name="weight",
89
+ )
90
+ if self.use_bias:
91
+ self.bias = self.add_weight(
92
+ shape=(self.out_channels,),
93
+ initializer=b_init,
94
+ trainable=True,
95
+ name="bias",
96
+ )
97
+ else:
98
+ self.bias = None
99
+
100
+ def build(self, input_shape):
101
+ if not hasattr(self, 'weight') or self.weight is None:
102
+ self._build_weights(input_shape[-1])
103
+ super().build(input_shape)
104
+
105
+ def reset_parameters(self):
106
+ if hasattr(self, 'weight') and self.weight is not None:
107
+ w_init = _get_weight_initializer(self.weight_initializer, self.in_channels)
108
+ self.weight.assign(w_init(self.weight.shape, dtype=self.weight.dtype))
109
+ if hasattr(self, 'bias') and self.bias is not None:
110
+ b_init = _get_bias_initializer(self.bias_initializer, self.in_channels)
111
+ self.bias.assign(b_init(self.bias.shape, dtype=self.bias.dtype))
112
+
113
+ def call(self, x):
114
+ if not hasattr(self, 'weight') or self.weight is None:
115
+ self._build_weights(ops.shape(x)[-1])
116
+
117
+ out = ops.matmul(x, self.weight)
118
+ if self.bias is not None:
119
+ out = out + self.bias
120
+ return out
121
+
122
+ def compute_output_shape(self, input_shape):
123
+ return (*input_shape[:-1], self.out_channels)
124
+
125
+ def __repr__(self) -> str:
126
+ return (f'{self.__class__.__name__}({self.in_channels}, '
127
+ f'{self.out_channels}, bias={self.use_bias})')
128
+
129
+
130
+ class HeteroLinear(layers.Layer):
131
+ r"""Applies separate linear transformations to the incoming data according
132
+ to types.
133
+
134
+ .. math::
135
+ \mathbf{x}^{\prime}_i = \mathbf{x}_i \mathbf{\Theta}_{\kappa_i} +
136
+ \mathbf{b}_{\kappa_i}
137
+
138
+ Args:
139
+ in_channels (int): Size of each input sample.
140
+ out_channels (int): Size of each output sample.
141
+ num_types (int): The number of types.
142
+ is_sorted (bool, optional): If set to :obj:`True`, assumes that
143
+ :obj:`type_vec` is sorted. (default: :obj:`False`)
144
+ bias (bool, optional): If set to :obj:`False`, the layer will not learn
145
+ an additive bias. (default: :obj:`True`)
146
+ weight_initializer (str, optional): The initializer for the weight
147
+ matrix (:obj:`"glorot"`, :obj:`"uniform"`, :obj:`"kaiming_uniform"`
148
+ or :obj:`None`). (default: :obj:`None`)
149
+ bias_initializer (str, optional): The initializer for the bias vector
150
+ (:obj:`"zeros"` or :obj:`None`). (default: :obj:`None`)
151
+
152
+ Example:
153
+ ```python
154
+ import numpy as np
155
+ from k3_node.layers import HeteroLinear
156
+
157
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
158
+ node_type = np.random.randint(0, 3, size=(10,))
159
+
160
+ layer = HeteroLinear(in_channels=8, out_channels=16, num_types=3) # separate weights per type
161
+ out = layer(x, node_type)
162
+ print(tuple(out.shape)) # (10, 16)
163
+ ```
164
+ """
165
+ def __init__(
166
+ self,
167
+ in_channels: int,
168
+ out_channels: int,
169
+ num_types: int,
170
+ is_sorted: bool = False,
171
+ bias: bool = True,
172
+ weight_initializer: Optional[str] = None,
173
+ bias_initializer: Optional[str] = None,
174
+ **kwargs
175
+ ):
176
+ super().__init__(**kwargs)
177
+ self.in_channels = in_channels
178
+ self.out_channels = out_channels
179
+ self.num_types = num_types
180
+ self.is_sorted = is_sorted
181
+ self.use_bias = bias
182
+ self.weight_initializer = weight_initializer
183
+ self.bias_initializer = bias_initializer
184
+ self.weight = None
185
+ self.bias = None
186
+
187
+ if self.in_channels > 0:
188
+ self._build_weights(self.in_channels)
189
+
190
+ def _build_weights(self, in_features: int):
191
+ self.in_channels = in_features
192
+ w_init = _get_weight_initializer(self.weight_initializer, in_features)
193
+ b_init = _get_bias_initializer(self.bias_initializer, in_features)
194
+
195
+ self.weight = self.add_weight(
196
+ shape=(self.num_types, in_features, self.out_channels),
197
+ initializer=w_init,
198
+ trainable=True,
199
+ name="weight",
200
+ )
201
+ if self.use_bias:
202
+ self.bias = self.add_weight(
203
+ shape=(self.num_types, self.out_channels),
204
+ initializer=b_init,
205
+ trainable=True,
206
+ name="bias",
207
+ )
208
+ else:
209
+ self.bias = None
210
+
211
+ def build(self, input_shape):
212
+ if self.in_channels <= 0:
213
+ self._build_weights(input_shape[-1])
214
+ super().build(input_shape)
215
+
216
+ def reset_parameters(self):
217
+ if self.in_channels > 0 and self.weight is not None:
218
+ w_init = _get_weight_initializer(self.weight_initializer, self.in_channels)
219
+ self.weight.assign(w_init(self.weight.shape, dtype=self.weight.dtype))
220
+ if self.bias is not None:
221
+ b_init = _get_bias_initializer(self.bias_initializer, self.in_channels)
222
+ self.bias.assign(b_init(self.bias.shape, dtype=self.bias.dtype))
223
+
224
+ def call(self, x, type_vec):
225
+ if self.in_channels <= 0 and not self.built:
226
+ self.build(ops.shape(x))
227
+
228
+ num_nodes = ops.shape(x)[0]
229
+ type_vec = ops.cast(type_vec, dtype="int32")
230
+
231
+ w_selected = ops.take(self.weight, type_vec, axis=0)
232
+ x_exp = ops.expand_dims(x, axis=1)
233
+ out = ops.squeeze(ops.matmul(x_exp, w_selected), axis=1)
234
+
235
+ if self.bias is not None:
236
+ b_selected = ops.take(self.bias, type_vec, axis=0)
237
+ out = out + b_selected
238
+
239
+ return out
240
+
241
+ def compute_output_shape(self, input_shape):
242
+ if isinstance(input_shape, (tuple, list)):
243
+ x_shape = input_shape[0]
244
+ else:
245
+ x_shape = input_shape
246
+ return (*x_shape[:-1], self.out_channels)
247
+
248
+ def __repr__(self) -> str:
249
+ return (f'{self.__class__.__name__}({self.in_channels}, '
250
+ f'{self.out_channels}, num_types={self.num_types}, '
251
+ f'bias={self.use_bias})')
252
+
253
+
254
+ class HeteroDictLinear(layers.Layer):
255
+ r"""Applies separate linear transformations to the incoming data
256
+ dictionary.
257
+
258
+ For key :math:`\kappa`, it computes
259
+
260
+ .. math::
261
+ \mathbf{x}^{\prime}_{\kappa} = \mathbf{x}_{\kappa}
262
+ \mathbf{W}_{\kappa} + \mathbf{b}_{\kappa}.
263
+
264
+ Args:
265
+ in_channels (int or Dict[Any, int]): Size of each input sample.
266
+ out_channels (int): Size of each output sample.
267
+ types (List[Any], optional): The keys of the input dictionary.
268
+ (default: :obj:`None`)
269
+
270
+ Example:
271
+ ```python
272
+ import numpy as np
273
+ from k3_node.layers import HeteroDictLinear
274
+
275
+ x_dict = {"author": np.random.rand(3, 8).astype("float32"), "paper": np.random.rand(4, 12).astype("float32")}
276
+ layer = HeteroDictLinear(in_channels={"author": 8, "paper": 12}, out_channels=16)
277
+ out_dict = layer(x_dict)
278
+ print(tuple(out_dict["author"].shape), tuple(out_dict["paper"].shape)) # (3, 16) (4, 16)
279
+ ```
280
+ """
281
+ def __init__(
282
+ self,
283
+ in_channels: Union[int, Dict[Any, int]],
284
+ out_channels: int,
285
+ types: Optional[List[Any]] = None,
286
+ **kwargs
287
+ ):
288
+ layer_kwargs = {k: v for k, v in kwargs.items() if k in ["name", "trainable", "dtype", "autocast"]}
289
+ super().__init__(**layer_kwargs)
290
+
291
+ if isinstance(in_channels, dict):
292
+ self.types = list(in_channels.keys())
293
+ if types is not None and set(self.types) != set(types):
294
+ raise ValueError("The provided 'types' do not match with the "
295
+ "keys in the 'in_channels' dictionary")
296
+ else:
297
+ if types is None:
298
+ raise ValueError("Please provide a list of 'types' if passing "
299
+ "'in_channels' as an integer")
300
+ self.types = types
301
+ in_channels = {node_type: in_channels for node_type in types}
302
+
303
+ self.in_channels = in_channels
304
+ self.out_channels = out_channels
305
+ self.kwargs = kwargs
306
+
307
+ lin_kwargs = {k: v for k, v in kwargs.items() if k not in ["name"]}
308
+ self.lins = {
309
+ str(key): Linear(channels, self.out_channels, name=f"lin_{key}", **lin_kwargs)
310
+ for key, channels in self.in_channels.items()
311
+ }
312
+
313
+ def reset_parameters(self):
314
+ for lin in self.lins.values():
315
+ lin.reset_parameters()
316
+
317
+ def call(self, x_dict: Dict[str, Any]) -> Dict[str, Any]:
318
+ out_dict = {}
319
+ for key, x in x_dict.items():
320
+ str_key = str(key)
321
+ if str_key in self.lins:
322
+ out_dict[key] = self.lins[str_key](x)
323
+ return out_dict
324
+
325
+ def __repr__(self) -> str:
326
+ return (f'{self.__class__.__name__}({self.in_channels}, '
327
+ f'{self.out_channels}, bias={self.kwargs.get("bias", True)})')
@@ -0,0 +1,92 @@
1
+ from typing import Optional, Tuple
2
+ from keras import ops
3
+
4
+
5
+ def dense_mincut_pool(
6
+ x,
7
+ adj,
8
+ s,
9
+ mask: Optional[any] = None,
10
+ temp: float = 1.0,
11
+ ) -> Tuple[any, any, any, any]:
12
+ r"""The MinCut pooling operator from the `"Spectral Clustering in Graph
13
+ Neural Networks for Graph Pooling" <https://arxiv.org/abs/1907.00481>`_
14
+ paper.
15
+
16
+ .. math::
17
+ \mathbf{X}^{\prime} &= {\mathrm{softmax}(\mathbf{S})}^{\top} \cdot
18
+ \mathbf{X}
19
+
20
+ \mathbf{A}^{\prime} &= {\mathrm{softmax}(\mathbf{S})}^{\top} \cdot
21
+ \mathbf{A} \cdot \mathrm{softmax}(\mathbf{S})
22
+
23
+ Args:
24
+ x: Node feature tensor [B, N, F] or [N, F].
25
+ adj: Adjacency tensor [B, N, N] or [N, N].
26
+ s: Assignment tensor [B, N, C] or [N, C].
27
+ mask: Mask tensor [B, N] indicating valid nodes. (default: None)
28
+ temp: Temperature parameter for softmax function. (default: 1.0)
29
+
30
+ Example:
31
+ ```python
32
+ import numpy as np
33
+ from k3_node.layers import dense_mincut_pool
34
+
35
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
36
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
37
+ s = np.random.rand(2, 10, 3).astype("float32") # assignment scores for 3 clusters
38
+
39
+ x_pool, adj_pool, mincut_loss, ortho_loss = dense_mincut_pool(x, adj, s)
40
+ print(tuple(x_pool.shape), tuple(adj_pool.shape)) # (2, 3, 8) (2, 3, 3)
41
+ ```
42
+ """
43
+ if len(ops.shape(x)) == 2:
44
+ x = ops.expand_dims(x, axis=0)
45
+ if len(ops.shape(adj)) == 2:
46
+ adj = ops.expand_dims(adj, axis=0)
47
+ if len(ops.shape(s)) == 2:
48
+ s = ops.expand_dims(s, axis=0)
49
+
50
+ batch_size = ops.shape(x)[0]
51
+ num_nodes = ops.shape(x)[1]
52
+ k = ops.shape(s)[-1]
53
+
54
+ if temp != 1.0:
55
+ s = s / temp
56
+ s = ops.softmax(s, axis=-1)
57
+
58
+ if mask is not None:
59
+ mask_m = ops.cast(ops.reshape(mask, (batch_size, num_nodes, 1)), x.dtype)
60
+ x = x * mask_m
61
+ s = s * mask_m
62
+
63
+ s_t = ops.transpose(s, (0, 2, 1))
64
+
65
+ out = ops.matmul(s_t, x)
66
+ out_adj = ops.matmul(ops.matmul(s_t, adj), s)
67
+
68
+ # MinCut regularization
69
+ mincut_num = ops.sum(ops.diagonal(out_adj, axis1=1, axis2=2), axis=-1)
70
+ d_flat = ops.sum(adj, axis=-1)
71
+ d = ops.expand_dims(d_flat, axis=-1) * ops.expand_dims(ops.eye(num_nodes, dtype=d_flat.dtype), axis=0)
72
+ s_d_s = ops.matmul(ops.matmul(s_t, d), s)
73
+ mincut_den = ops.sum(ops.diagonal(s_d_s, axis1=1, axis2=2), axis=-1)
74
+ mincut_loss = ops.mean(-(mincut_num / mincut_den))
75
+
76
+ # Orthogonality regularization
77
+ ss = ops.matmul(s_t, s)
78
+ norm_ss = ops.sqrt(ops.sum(ops.power(ss, 2), axis=(-1, -2), keepdims=True))
79
+ i_s = ops.eye(k, dtype=ss.dtype)
80
+ norm_is = ops.sqrt(ops.cast(k, ss.dtype))
81
+ diff = (ss / norm_ss) - (ops.expand_dims(i_s, axis=0) / norm_is)
82
+ ortho_loss = ops.mean(ops.sqrt(ops.sum(ops.power(diff, 2), axis=(-1, -2))))
83
+
84
+ # Fix and normalize coarsened adjacency matrix
85
+ eye_k = ops.expand_dims(ops.eye(k, dtype=out_adj.dtype), axis=0)
86
+ out_adj = out_adj * (1.0 - eye_k)
87
+ d_out = ops.sum(out_adj, axis=-1, keepdims=True)
88
+ d_sqrt = ops.sqrt(d_out) + 1e-15
89
+ out_adj = (out_adj / d_sqrt) / ops.transpose(d_sqrt, (0, 2, 1))
90
+
91
+ return out, out_adj, mincut_loss, ortho_loss
92
+