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,412 @@
1
+ from typing import Optional
2
+ from keras import initializers, ops
3
+
4
+ from .base import Aggregation
5
+ from k3_node.ops.segment import segment_max, segment_sum
6
+ from k3_node.ops.creation import full
7
+
8
+
9
+ class SumAggregation(Aggregation):
10
+ r"""An aggregation operator that sums up features across a set of elements.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import SumAggregation
16
+
17
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
18
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
19
+
20
+ aggr = SumAggregation()
21
+ out = aggr(x, index=index, dim_size=2)
22
+ print(tuple(out.shape)) # (2, 8)
23
+ ```
24
+ """
25
+
26
+ def call(
27
+ self,
28
+ x,
29
+ index: Optional[any] = None,
30
+ ptr: Optional[any] = None,
31
+ dim_size: Optional[int] = None,
32
+ dim: int = -2,
33
+ **kwargs,
34
+ ):
35
+ return self.reduce(x, index, ptr, dim_size, dim, reduce="sum")
36
+
37
+
38
+ class MeanAggregation(Aggregation):
39
+ r"""An aggregation operator that averages features across a set of elements.
40
+
41
+ Example:
42
+ ```python
43
+ import numpy as np
44
+ from k3_node.layers import MeanAggregation
45
+
46
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
47
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
48
+
49
+ aggr = MeanAggregation()
50
+ out = aggr(x, index=index, dim_size=2)
51
+ print(tuple(out.shape)) # (2, 8)
52
+ ```
53
+ """
54
+
55
+ def call(
56
+ self,
57
+ x,
58
+ index: Optional[any] = None,
59
+ ptr: Optional[any] = None,
60
+ dim_size: Optional[int] = None,
61
+ dim: int = -2,
62
+ **kwargs,
63
+ ):
64
+ return self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
65
+
66
+
67
+ class MaxAggregation(Aggregation):
68
+ r"""An aggregation operator that takes the feature-wise maximum across a set of elements.
69
+
70
+ Example:
71
+ ```python
72
+ import numpy as np
73
+ from k3_node.layers import MaxAggregation
74
+
75
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
76
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
77
+
78
+ aggr = MaxAggregation()
79
+ out = aggr(x, index=index, dim_size=2)
80
+ print(tuple(out.shape)) # (2, 8)
81
+ ```
82
+ """
83
+
84
+ def call(
85
+ self,
86
+ x,
87
+ index: Optional[any] = None,
88
+ ptr: Optional[any] = None,
89
+ dim_size: Optional[int] = None,
90
+ dim: int = -2,
91
+ **kwargs,
92
+ ):
93
+ return self.reduce(x, index, ptr, dim_size, dim, reduce="max")
94
+
95
+
96
+ class MinAggregation(Aggregation):
97
+ r"""An aggregation operator that takes the feature-wise minimum across a set of elements.
98
+
99
+ Example:
100
+ ```python
101
+ import numpy as np
102
+ from k3_node.layers import MinAggregation
103
+
104
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
105
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
106
+
107
+ aggr = MinAggregation()
108
+ out = aggr(x, index=index, dim_size=2)
109
+ print(tuple(out.shape)) # (2, 8)
110
+ ```
111
+ """
112
+
113
+ def call(
114
+ self,
115
+ x,
116
+ index: Optional[any] = None,
117
+ ptr: Optional[any] = None,
118
+ dim_size: Optional[int] = None,
119
+ dim: int = -2,
120
+ **kwargs,
121
+ ):
122
+ return self.reduce(x, index, ptr, dim_size, dim, reduce="min")
123
+
124
+
125
+ class MulAggregation(Aggregation):
126
+ r"""An aggregation operator that multiplies features across a set of elements.
127
+
128
+ Example:
129
+ ```python
130
+ import numpy as np
131
+ from k3_node.layers import MulAggregation
132
+
133
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
134
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
135
+
136
+ aggr = MulAggregation()
137
+ out = aggr(x, index=index, dim_size=2)
138
+ print(tuple(out.shape)) # (2, 8)
139
+ ```
140
+ """
141
+
142
+ def call(
143
+ self,
144
+ x,
145
+ index: Optional[any] = None,
146
+ ptr: Optional[any] = None,
147
+ dim_size: Optional[int] = None,
148
+ dim: int = -2,
149
+ **kwargs,
150
+ ):
151
+ self.assert_index_present(index)
152
+ return self.reduce(x, index, ptr, dim_size, dim, reduce="mul")
153
+
154
+
155
+ class VarAggregation(Aggregation):
156
+ r"""An aggregation operator that takes the feature-wise variance across a set of elements.
157
+
158
+ Example:
159
+ ```python
160
+ import numpy as np
161
+ from k3_node.layers import VarAggregation
162
+
163
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
164
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
165
+
166
+ aggr = VarAggregation()
167
+ out = aggr(x, index=index, dim_size=2)
168
+ print(tuple(out.shape)) # (2, 8)
169
+ ```
170
+ """
171
+
172
+ def __init__(self, semi_grad: bool = False, **kwargs):
173
+ super().__init__(**kwargs)
174
+ self.semi_grad = semi_grad
175
+
176
+ def call(
177
+ self,
178
+ x,
179
+ index: Optional[any] = None,
180
+ ptr: Optional[any] = None,
181
+ dim_size: Optional[int] = None,
182
+ dim: int = -2,
183
+ **kwargs,
184
+ ):
185
+ mean = self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
186
+ x_sq = ops.power(x, 2)
187
+ if self.semi_grad:
188
+ x_sq = ops.stop_gradient(x_sq)
189
+ mean2 = self.reduce(x_sq, index, ptr, dim_size, dim, reduce="mean")
190
+ return mean2 - ops.power(mean, 2)
191
+
192
+ def __repr__(self) -> str:
193
+ return f"{self.__class__.__name__}(semi_grad={self.semi_grad})"
194
+
195
+
196
+ class StdAggregation(Aggregation):
197
+ r"""An aggregation operator that takes the feature-wise standard deviation across a set of elements.
198
+
199
+ Example:
200
+ ```python
201
+ import numpy as np
202
+ from k3_node.layers import StdAggregation
203
+
204
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
205
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
206
+
207
+ aggr = StdAggregation()
208
+ out = aggr(x, index=index, dim_size=2)
209
+ print(tuple(out.shape)) # (2, 8)
210
+ ```
211
+ """
212
+
213
+ def __init__(self, semi_grad: bool = False, **kwargs):
214
+ super().__init__(**kwargs)
215
+ self.semi_grad = semi_grad
216
+ self.var_aggr = VarAggregation(semi_grad=semi_grad)
217
+
218
+ def call(
219
+ self,
220
+ x,
221
+ index: Optional[any] = None,
222
+ ptr: Optional[any] = None,
223
+ dim_size: Optional[int] = None,
224
+ dim: int = -2,
225
+ **kwargs,
226
+ ):
227
+ var = self.var_aggr(x, index, ptr, dim_size, dim)
228
+ out = ops.sqrt(ops.maximum(var, 1e-5))
229
+ out = ops.where(out <= (1e-5**0.5), 0.0, out)
230
+ return out
231
+
232
+ def __repr__(self) -> str:
233
+ return f"{self.__class__.__name__}(semi_grad={self.semi_grad})"
234
+
235
+
236
+ class SoftmaxAggregation(Aggregation):
237
+ r"""The softmax aggregation operator based on a temperature term.
238
+
239
+ Example:
240
+ ```python
241
+ import numpy as np
242
+ from k3_node.layers import SoftmaxAggregation
243
+
244
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
245
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
246
+
247
+ aggr = SoftmaxAggregation(learn=True)
248
+ out = aggr(x, index=index, dim_size=2)
249
+ print(tuple(out.shape)) # (2, 8)
250
+ ```
251
+ """
252
+
253
+ def __init__(
254
+ self,
255
+ t: float = 1.0,
256
+ learn: bool = False,
257
+ semi_grad: bool = False,
258
+ channels: int = 1,
259
+ **kwargs,
260
+ ):
261
+ super().__init__(**kwargs)
262
+
263
+ if learn and semi_grad:
264
+ raise ValueError(
265
+ f"Cannot enable 'semi_grad' in '{self.__class__.__name__}' in "
266
+ f"case the temperature term 't' is learnable"
267
+ )
268
+
269
+ if not learn and channels != 1:
270
+ raise ValueError(
271
+ f"Cannot set 'channels' greater than '1' in case "
272
+ f"'{self.__class__.__name__}' is not trainable"
273
+ )
274
+
275
+ self._init_t = t
276
+ self.learn = learn
277
+ self.semi_grad = semi_grad
278
+ self.channels = channels
279
+
280
+ if learn:
281
+ self.t = self.add_weight(
282
+ shape=(channels,),
283
+ initializer=initializers.Constant(t),
284
+ trainable=True,
285
+ name="t",
286
+ )
287
+ else:
288
+ self.t = t
289
+
290
+ def reset_parameters(self):
291
+ if self.learn:
292
+ self.t.assign(full((self.channels,), self._init_t, dtype=self.t.dtype))
293
+
294
+ def call(
295
+ self,
296
+ x,
297
+ index: Optional[any] = None,
298
+ ptr: Optional[any] = None,
299
+ dim_size: Optional[int] = None,
300
+ dim: int = -2,
301
+ **kwargs,
302
+ ):
303
+ t = self.t
304
+ if self.channels != 1:
305
+ self.assert_two_dimensional_input(x, dim)
306
+ t = ops.reshape(t, (1, self.channels))
307
+
308
+ alpha = x
309
+ if self.learn or t != 1.0:
310
+ alpha = x * t
311
+
312
+ if not self.learn and self.semi_grad:
313
+ alpha = ops.stop_gradient(alpha)
314
+
315
+ # Graph-wise softmax over segments
316
+ index = ops.cast(index, dtype="int32")
317
+ if dim_size is None: # a tensor while tracing; don't test its truth value
318
+ dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
319
+
320
+ max_val = segment_max(alpha, index, num_segments=dim_size)
321
+ max_exp = ops.take(max_val, index, axis=0)
322
+ exp_alpha = ops.exp(alpha - max_exp)
323
+ sum_exp = segment_sum(exp_alpha, index, num_segments=dim_size)
324
+ sum_exp_taken = ops.take(sum_exp, index, axis=0)
325
+ alpha_sm = exp_alpha / ops.maximum(sum_exp_taken, 1e-12)
326
+
327
+ return self.reduce(x * alpha_sm, index, ptr, dim_size, dim, reduce="sum")
328
+
329
+ def __repr__(self) -> str:
330
+ return f"{self.__class__.__name__}(learn={self.learn})"
331
+
332
+
333
+ class PowerMeanAggregation(Aggregation):
334
+ r"""The powermean aggregation operator based on a power term.
335
+
336
+ Example:
337
+ ```python
338
+ import numpy as np
339
+ from k3_node.layers import PowerMeanAggregation
340
+
341
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
342
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
343
+
344
+ aggr = PowerMeanAggregation(learn=True)
345
+ out = aggr(x, index=index, dim_size=2)
346
+ print(tuple(out.shape)) # (2, 8)
347
+ ```
348
+ """
349
+
350
+ def __init__(
351
+ self,
352
+ p: float = 1.0,
353
+ learn: bool = False,
354
+ channels: int = 1,
355
+ clamp_min: Optional[float] = 1e-4,
356
+ clamp_max: Optional[float] = 100.0,
357
+ **kwargs,
358
+ ):
359
+ super().__init__(**kwargs)
360
+
361
+ if not learn and channels != 1:
362
+ raise ValueError(
363
+ f"Cannot set 'channels' greater than '1' in case "
364
+ f"'{self.__class__.__name__}' is not trainable"
365
+ )
366
+
367
+ self._init_p = p
368
+ self.learn = learn
369
+ self.channels = channels
370
+ self.min_value = clamp_min if clamp_min is not None else 1e-4
371
+ self.max_value = clamp_max if clamp_max is not None else 100.0
372
+
373
+ if learn:
374
+ self.p = self.add_weight(
375
+ shape=(channels,),
376
+ initializer=initializers.Constant(p),
377
+ trainable=True,
378
+ name="p",
379
+ )
380
+ else:
381
+ self.p = p
382
+
383
+ def reset_parameters(self):
384
+ if self.learn:
385
+ self.p.assign(full((self.channels,), self._init_p, dtype=self.p.dtype))
386
+
387
+ def call(
388
+ self,
389
+ x,
390
+ index: Optional[any] = None,
391
+ ptr: Optional[any] = None,
392
+ dim_size: Optional[int] = None,
393
+ dim: int = -2,
394
+ **kwargs,
395
+ ):
396
+ p = self.p
397
+ if self.channels != 1:
398
+ self.assert_two_dimensional_input(x, dim)
399
+ p = ops.reshape(p, (-1, self.channels))
400
+
401
+ if self.learn or p != 1.0:
402
+ x = ops.power(ops.clip(x, self.min_value, self.max_value), p)
403
+
404
+ out = self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
405
+
406
+ if self.learn or p != 1.0:
407
+ out = ops.power(ops.clip(out, self.min_value, self.max_value), 1.0 / p)
408
+
409
+ return out
410
+
411
+ def __repr__(self) -> str:
412
+ return f"{self.__class__.__name__}(learn={self.learn})"
@@ -0,0 +1,65 @@
1
+ from typing import Optional
2
+ from .base import Aggregation
3
+
4
+
5
+ class DeepSetsAggregation(Aggregation):
6
+ r"""Performs Deep Sets aggregation in which the elements to aggregate are
7
+ first transformed by a Multi-Layer Perceptron (MLP)
8
+ :math:`\phi_{\mathbf{\Theta}}`, summed, and then transformed by another MLP
9
+ :math:`\rho_{\mathbf{\Theta}}`.
10
+
11
+ Example:
12
+ ```python
13
+ import numpy as np
14
+ import keras
15
+ from k3_node.layers import DeepSetsAggregation
16
+
17
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
18
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
19
+
20
+ aggr = DeepSetsAggregation(local_nn=keras.layers.Dense(16), global_nn=keras.layers.Dense(16))
21
+ out = aggr(x, index=index, dim_size=2)
22
+ print(tuple(out.shape)) # (2, 16)
23
+ ```
24
+ """
25
+
26
+ def __init__(
27
+ self,
28
+ local_nn: Optional[any] = None,
29
+ global_nn: Optional[any] = None,
30
+ local_mlp: Optional[any] = None,
31
+ global_mlp: Optional[any] = None,
32
+ **kwargs,
33
+ ):
34
+ super().__init__(**kwargs)
35
+ self.local_nn = local_nn if local_nn is not None else local_mlp
36
+ self.global_nn = global_nn if global_nn is not None else global_mlp
37
+
38
+ def reset_parameters(self):
39
+ if self.local_nn is not None and hasattr(self.local_nn, "reset_parameters"):
40
+ self.local_nn.reset_parameters()
41
+ if self.global_nn is not None and hasattr(self.global_nn, "reset_parameters"):
42
+ self.global_nn.reset_parameters()
43
+
44
+ def call(
45
+ self,
46
+ x,
47
+ index: Optional[any] = None,
48
+ ptr: Optional[any] = None,
49
+ dim_size: Optional[int] = None,
50
+ dim: int = -2,
51
+ **kwargs,
52
+ ):
53
+ if self.local_nn is not None:
54
+ x = self.local_nn(x)
55
+
56
+ x = self.reduce(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim, reduce="sum")
57
+
58
+ if self.global_nn is not None:
59
+ x = self.global_nn(x)
60
+
61
+ return x
62
+
63
+ def __repr__(self) -> str:
64
+ return f"{self.__class__.__name__}(local_nn={self.local_nn}, global_nn={self.global_nn})"
65
+
@@ -0,0 +1,29 @@
1
+ from keras import ops
2
+
3
+ from k3_node.layers.aggr import Aggregation
4
+
5
+
6
+
7
+ class DeepSetsAggregation(Aggregation):
8
+ def __init__(
9
+ self,
10
+ local_mlp=None,
11
+ global_mlp=None,
12
+ ):
13
+ super().__init__()
14
+
15
+ self.local_mlp = local_mlp
16
+ self.global_mlp = global_mlp
17
+
18
+
19
+ def call(self, x, index=None, axis=-2):
20
+
21
+ if self.local_mlp is not None:
22
+ x = self.local_mlp(x)
23
+
24
+ x = self.reduce(x, index=index, axis=axis, reduce_fn=ops.segment_sum)
25
+
26
+ if self.global_mlp is not None:
27
+ x = self.global_mlp(x) # Assuming batch handling within MLP
28
+
29
+ return x
@@ -0,0 +1,107 @@
1
+ from typing import List, Optional
2
+ from keras import layers, ops
3
+
4
+ from .base import Aggregation
5
+
6
+
7
+ class ResNetPotential(layers.Layer):
8
+ def __init__(self, in_channels: int, out_channels: int, num_layers: List[int], **kwargs):
9
+ super().__init__(**kwargs)
10
+ sizes = [in_channels] + num_layers + [out_channels]
11
+ self.layers_list = [
12
+ layers.Dense(out_size, activation="tanh")
13
+ for in_size, out_size in zip(sizes[:-2], sizes[1:-1])
14
+ ]
15
+ self.final_layer = layers.Dense(sizes[-1])
16
+ self.res_trans = [
17
+ layers.Dense(layer_size)
18
+ for layer_size in num_layers + [out_channels]
19
+ ]
20
+
21
+ def reset_parameters(self):
22
+ for l in self.layers_list:
23
+ l.reset_parameters()
24
+ self.final_layer.reset_parameters()
25
+ for r in self.res_trans:
26
+ r.reset_parameters()
27
+
28
+ def call(self, inp):
29
+ h = inp
30
+ for layer, res in zip(self.layers_list + [self.final_layer], self.res_trans):
31
+ h_next = layer(h)
32
+ h = res(inp) + h_next
33
+ return h
34
+
35
+
36
+ class EquilibriumAggregation(Aggregation):
37
+ r"""The equilibrium aggregation layer from the `"Equilibrium Aggregation:
38
+ Encoding Sets via Optimization" <https://arxiv.org/abs/2202.12795>`_ paper.
39
+
40
+ Example:
41
+ ```python
42
+ import numpy as np
43
+ from k3_node.layers import EquilibriumAggregation
44
+
45
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
46
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
47
+
48
+ aggr = EquilibriumAggregation(in_channels=8, out_channels=16, num_layers=[8])
49
+ out = aggr(x, index=index, dim_size=2)
50
+ print(tuple(out.shape)) # (2, 16)
51
+ ```
52
+ """
53
+
54
+ def __init__(
55
+ self,
56
+ in_channels: int,
57
+ out_channels: int,
58
+ num_layers: List[int],
59
+ grad_iter: int = 5,
60
+ lamb: float = 0.1,
61
+ **kwargs,
62
+ ):
63
+ super().__init__(**kwargs)
64
+ self.in_channels = in_channels
65
+ self.out_channels = out_channels
66
+ self.num_layers = num_layers
67
+ self.grad_iter = grad_iter
68
+ self.initial_lamb = lamb
69
+
70
+ self.potential = ResNetPotential(in_channels + out_channels, out_channels, num_layers)
71
+ self.proj = layers.Dense(out_channels)
72
+
73
+ def reset_parameters(self):
74
+ self.potential.reset_parameters()
75
+ self.proj.reset_parameters()
76
+
77
+ def call(
78
+ self,
79
+ x,
80
+ index: Optional[any] = None,
81
+ ptr: Optional[any] = None,
82
+ dim_size: Optional[int] = None,
83
+ dim: int = -2,
84
+ **kwargs,
85
+ ):
86
+ self.assert_index_present(index)
87
+ index = ops.cast(index, dtype="int32")
88
+ if dim_size is None:
89
+ dim_size = ops.max(index) + 1
90
+
91
+ # Initial mean aggregation as starting state
92
+ x_mean = self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
93
+ y = self.proj(x_mean)
94
+
95
+ # Unrolled iterative equilibrium updates
96
+ for _ in range(self.grad_iter):
97
+ y_expanded = ops.take(y, index, axis=0)
98
+ inp = ops.concatenate([x, y_expanded], axis=-1)
99
+ pot = self.potential(inp)
100
+ pot_mean = self.reduce(pot, index, ptr, dim_size, dim, reduce="mean")
101
+ y = y + 0.1 * pot_mean - self.initial_lamb * y
102
+
103
+ return y
104
+
105
+ def __repr__(self) -> str:
106
+ return f"{self.__class__.__name__}()"
107
+
@@ -0,0 +1,43 @@
1
+ from typing import List, Union
2
+ from .base import Aggregation
3
+
4
+
5
+ class FusedAggregation(Aggregation):
6
+ r"""Helper class to fuse computation of multiple aggregations together.
7
+
8
+ Example:
9
+ ```python
10
+ import numpy as np
11
+ from k3_node.layers import FusedAggregation
12
+
13
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
14
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
15
+
16
+ aggr = FusedAggregation(aggrs=["sum", "mean", "max"])
17
+ outs = aggr(x, index=index, dim_size=2) # one result per aggregation, computed together
18
+ print(len(outs), tuple(outs[0].shape)) # 3 (2, 8)
19
+ ```
20
+ """
21
+
22
+ def __init__(self, aggrs: List[Union[Aggregation, str]], **kwargs):
23
+ super().__init__(**kwargs)
24
+ from .resolver import aggregation_resolver
25
+
26
+ self.aggrs = [aggregation_resolver(aggr) for aggr in aggrs]
27
+
28
+ def reset_parameters(self):
29
+ for aggr in self.aggrs:
30
+ if hasattr(aggr, "reset_parameters"):
31
+ aggr.reset_parameters()
32
+
33
+ def call(
34
+ self,
35
+ x,
36
+ index=None,
37
+ ptr=None,
38
+ dim_size=None,
39
+ dim=-2,
40
+ **kwargs,
41
+ ):
42
+ return [aggr(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim, **kwargs) for aggr in self.aggrs]
43
+