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,89 @@
1
+ from typing import Optional
2
+ from keras import ops
3
+
4
+ from .base import Aggregation
5
+ from .utils import PoolingByMultiheadAttention, SetAttentionBlock
6
+
7
+
8
+ class GraphMultisetTransformer(Aggregation):
9
+ r"""The Graph Multiset Transformer pooling operator from the
10
+ `"Accurate Learning of Graph Representations
11
+ with Graph Multiset Pooling" <https://arxiv.org/abs/2102.11533>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import GraphMultisetTransformer
17
+
18
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
19
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
20
+
21
+ aggr = GraphMultisetTransformer(channels=8, k=2)
22
+ out = aggr(x, index=index, dim_size=2)
23
+ print(tuple(out.shape)) # (2, 8)
24
+ ```
25
+ """
26
+
27
+ def __init__(
28
+ self,
29
+ channels: int,
30
+ k: int,
31
+ num_encoder_blocks: int = 1,
32
+ heads: int = 1,
33
+ layer_norm: bool = False,
34
+ dropout: float = 0.0,
35
+ **kwargs,
36
+ ):
37
+ super().__init__(**kwargs)
38
+ self.channels = channels
39
+ self.k = k
40
+ self.heads = heads
41
+ self.layer_norm = layer_norm
42
+ self.dropout = dropout
43
+
44
+ self.pma1 = PoolingByMultiheadAttention(channels, k, heads, layer_norm, dropout)
45
+ self.encoders = [
46
+ SetAttentionBlock(channels, heads, layer_norm, dropout)
47
+ for _ in range(num_encoder_blocks)
48
+ ]
49
+ self.pma2 = PoolingByMultiheadAttention(channels, 1, heads, layer_norm, dropout)
50
+
51
+ def reset_parameters(self):
52
+ self.pma1.reset_parameters()
53
+ for encoder in self.encoders:
54
+ encoder.reset_parameters()
55
+ self.pma2.reset_parameters()
56
+
57
+ def call(
58
+ self,
59
+ x,
60
+ index: Optional[any] = None,
61
+ ptr: Optional[any] = None,
62
+ dim_size: Optional[int] = None,
63
+ dim: int = -2,
64
+ max_num_elements: Optional[int] = None,
65
+ training: bool = False,
66
+ **kwargs,
67
+ ):
68
+ x_dense, mask = self.to_dense_batch(
69
+ x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
70
+ max_num_elements=max_num_elements,
71
+ )
72
+
73
+ x_dense = self.pma1(x_dense, mask=mask, training=training)
74
+
75
+ for encoder in self.encoders:
76
+ x_dense = encoder(x_dense, training=training)
77
+
78
+ x_dense = self.pma2(x_dense, training=training)
79
+
80
+ # Output shape: [B, channels]
81
+ return ops.squeeze(x_dense, axis=1)
82
+
83
+ def __repr__(self) -> str:
84
+ return (
85
+ f"{self.__class__.__name__}({self.channels}, k={self.k}, "
86
+ f"heads={self.heads}, layer_norm={self.layer_norm}, "
87
+ f"dropout={self.dropout})"
88
+ )
89
+
@@ -0,0 +1,58 @@
1
+ from typing import Optional
2
+ from keras import layers
3
+
4
+ from .base import Aggregation
5
+
6
+
7
+ class GRUAggregation(Aggregation):
8
+ r"""Performs GRU aggregation in which the elements to aggregate are
9
+ interpreted as a sequence, as described in the `"Graph Neural Networks
10
+ with Adaptive Readouts" <https://arxiv.org/abs/2211.04952>`_ paper.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import GRUAggregation
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 = GRUAggregation(in_channels=8, out_channels=16)
21
+ out = aggr(x, index=index, dim_size=2)
22
+ print(tuple(out.shape)) # (2, 16)
23
+ ```
24
+ """
25
+
26
+ def __init__(self, in_channels: int, out_channels: int, **kwargs):
27
+ super().__init__(**kwargs)
28
+ self.in_channels = in_channels
29
+ self.out_channels = out_channels
30
+ self.gru = layers.GRU(out_channels, return_sequences=False)
31
+
32
+ def build(self, input_shape=None):
33
+ self.gru.build((None, None, self.in_channels))
34
+ super().build(input_shape)
35
+
36
+ def reset_parameters(self):
37
+ self.gru.reset_parameters()
38
+
39
+ def call(
40
+ self,
41
+ x,
42
+ index: Optional[any] = None,
43
+ ptr: Optional[any] = None,
44
+ dim_size: Optional[int] = None,
45
+ dim: int = -2,
46
+ max_num_elements: Optional[int] = None,
47
+ training: bool = False,
48
+ **kwargs,
49
+ ):
50
+ x_dense, _ = self.to_dense_batch(
51
+ x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
52
+ max_num_elements=max_num_elements,
53
+ )
54
+ return self.gru(x_dense, training=training)
55
+
56
+ def __repr__(self) -> str:
57
+ return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels})"
58
+
@@ -0,0 +1,143 @@
1
+ from math import ceil, log2
2
+ from typing import Optional
3
+ from keras import layers, ops
4
+
5
+ from .base import Aggregation
6
+
7
+
8
+ class LCMAggregation(Aggregation):
9
+ r"""The Learnable Commutative Monoid aggregation from the
10
+ `"Learnable Commutative Monoids for Graph Neural Networks"
11
+ <https://arxiv.org/abs/2212.08541>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import LCMAggregation
17
+
18
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
19
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
20
+
21
+ aggr = LCMAggregation(in_channels=8, out_channels=16)
22
+ out = aggr(x, index=index, dim_size=2)
23
+ print(tuple(out.shape)) # (2, 16)
24
+ ```
25
+ """
26
+
27
+ def __init__(
28
+ self,
29
+ in_channels: int,
30
+ out_channels: int,
31
+ project: bool = True,
32
+ **kwargs,
33
+ ):
34
+ super().__init__(**kwargs)
35
+ if in_channels != out_channels and not project:
36
+ raise ValueError(
37
+ f"Inputs of '{self.__class__.__name__}' must be projected if `in_channels != out_channels`"
38
+ )
39
+
40
+ self.in_channels = in_channels
41
+ self.out_channels = out_channels
42
+ self.project = project
43
+
44
+ self.lin = layers.Dense(out_channels) if project else None
45
+ self.gru_cell = layers.GRUCell(out_channels)
46
+
47
+ def reset_parameters(self):
48
+ if self.lin is not None:
49
+ self.lin.reset_parameters()
50
+ self.gru_cell.reset_parameters()
51
+
52
+ def call(
53
+ self,
54
+ x,
55
+ index: Optional[any] = None,
56
+ ptr: Optional[any] = None,
57
+ dim_size: Optional[int] = None,
58
+ dim: int = -2,
59
+ max_num_elements: Optional[int] = None,
60
+ **kwargs,
61
+ ):
62
+ if self.lin is not None:
63
+ x = ops.relu(self.lin(x))
64
+
65
+ x_dense, _ = self.to_dense_batch(
66
+ x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
67
+ max_num_elements=max_num_elements,
68
+ )
69
+
70
+ # Transpose to [num_neighbors, num_nodes, num_features]
71
+ x_dense = ops.transpose(x_dense, (1, 0, 2))
72
+ num_neighbors = ops.shape(x_dense)[0]
73
+ num_nodes = ops.shape(x_dense)[1]
74
+ num_features = ops.shape(x_dense)[2]
75
+
76
+ if not isinstance(num_neighbors, int): # a tensor while tracing (e.g. TensorFlow's fit)
77
+ return self._reduce_traced(x_dense)
78
+ if num_neighbors == 0:
79
+ return ops.zeros((num_nodes, self.out_channels), dtype=x.dtype)
80
+
81
+ depth = ceil(log2(max(num_neighbors, 1)))
82
+ for _ in range(depth):
83
+ curr_len = ops.shape(x_dense)[0]
84
+ if curr_len <= 1:
85
+ break
86
+ half_size = ceil(curr_len / 2)
87
+
88
+ if curr_len % 2 == 1:
89
+ x_pair = x_dense[:-1]
90
+ remainder = x_dense[-1:]
91
+ else:
92
+ x_pair = x_dense
93
+ remainder = None
94
+
95
+ # x_pair: [2 * half, num_nodes, num_features]
96
+ pair_count = ops.shape(x_pair)[0] // 2
97
+ x_pair = ops.reshape(x_pair, (pair_count, 2, num_nodes, num_features))
98
+ left = x_pair[:, 0] # [pair_count, num_nodes, num_features]
99
+ right = x_pair[:, 1] # [pair_count, num_nodes, num_features]
100
+
101
+ left_flat = ops.reshape(left, (-1, num_features))
102
+ right_flat = ops.reshape(right, (-1, num_features))
103
+
104
+ # GRUCell: inputs=left, state=[right]
105
+ out1, _ = self.gru_cell(left_flat, [right_flat])
106
+ out2, _ = self.gru_cell(right_flat, [left_flat])
107
+ out = 0.5 * (out1 + out2)
108
+ out = ops.reshape(out, (pair_count, num_nodes, num_features))
109
+
110
+ if remainder is not None:
111
+ out = ops.concatenate([out, remainder], axis=0)
112
+
113
+ x_dense = out
114
+
115
+ return ops.squeeze(x_dense, axis=0)
116
+
117
+ def _combine(self, left, right):
118
+ num_features = ops.shape(left)[-1]
119
+ left_flat, right_flat = ops.reshape(left, (-1, num_features)), ops.reshape(right, (-1, num_features))
120
+ out1, _ = self.gru_cell(left_flat, [right_flat])
121
+ out2, _ = self.gru_cell(right_flat, [left_flat])
122
+ return ops.reshape(0.5 * (out1 + out2), ops.shape(left))
123
+
124
+ def _reduce_traced(self, x_dense):
125
+ """The pairwise reduction of ``call`` for a neighbor count only known at run time: position
126
+ ``i`` absorbs position ``i + stride`` when ``i`` is a multiple of ``2 * stride``, with the
127
+ stride doubling every level. This pairs the same elements as ``call``, in a fixed shape."""
128
+ length = ops.shape(x_dense)[0]
129
+ if not self.gru_cell.built: # no weights may be created inside the loop
130
+ self.gru_cell.build((None, x_dense.shape[-1]))
131
+ position = ops.arange(length, dtype="int32")
132
+
133
+ def body(stride, h):
134
+ combined = self._combine(h, ops.roll(h, -stride, axis=0))
135
+ active = ops.logical_and(position % (2 * stride) == 0, position + stride < length)
136
+ return stride * 2, ops.where(active[:, None, None], combined, h)
137
+
138
+ _, h = ops.while_loop(lambda stride, h: stride < length, body, (ops.convert_to_tensor(1, "int32"), x_dense))
139
+ return h[0]
140
+
141
+ def __repr__(self) -> str:
142
+ return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels}, project={self.project})"
143
+
@@ -0,0 +1,58 @@
1
+ from typing import Optional
2
+ from keras import layers
3
+
4
+ from .base import Aggregation
5
+
6
+
7
+ class LSTMAggregation(Aggregation):
8
+ r"""Performs LSTM-style aggregation in which the elements to aggregate are
9
+ interpreted as a sequence, as described in the `"Inductive Representation
10
+ Learning on Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import LSTMAggregation
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 = LSTMAggregation(in_channels=8, out_channels=16)
21
+ out = aggr(x, index=index, dim_size=2)
22
+ print(tuple(out.shape)) # (2, 16)
23
+ ```
24
+ """
25
+
26
+ def __init__(self, in_channels: int, out_channels: int, **kwargs):
27
+ super().__init__(**kwargs)
28
+ self.in_channels = in_channels
29
+ self.out_channels = out_channels
30
+ self.lstm = layers.LSTM(out_channels, return_sequences=False)
31
+
32
+ def build(self, input_shape=None):
33
+ self.lstm.build((None, None, self.in_channels))
34
+ super().build(input_shape)
35
+
36
+ def reset_parameters(self):
37
+ self.lstm.reset_parameters()
38
+
39
+ def call(
40
+ self,
41
+ x,
42
+ index: Optional[any] = None,
43
+ ptr: Optional[any] = None,
44
+ dim_size: Optional[int] = None,
45
+ dim: int = -2,
46
+ max_num_elements: Optional[int] = None,
47
+ training: bool = False,
48
+ **kwargs,
49
+ ):
50
+ x_dense, _ = self.to_dense_batch(
51
+ x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
52
+ max_num_elements=max_num_elements,
53
+ )
54
+ return self.lstm(x_dense, training=training)
55
+
56
+ def __repr__(self) -> str:
57
+ return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels})"
58
+
@@ -0,0 +1,75 @@
1
+ from typing import Optional
2
+ from keras import layers, ops
3
+
4
+ from .base import Aggregation
5
+
6
+
7
+ class MLPAggregation(Aggregation):
8
+ r"""Performs MLP aggregation in which the elements to aggregate are
9
+ flattened into a single vectorial representation, and are then processed by
10
+ a Multi-Layer Perceptron (MLP).
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import MLPAggregation
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 = MLPAggregation(in_channels=8, out_channels=16, max_num_elements=5)
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
+ in_channels: int,
29
+ out_channels: int,
30
+ max_num_elements: int,
31
+ mlp: Optional[any] = None,
32
+ **kwargs,
33
+ ):
34
+ super().__init__(**kwargs)
35
+ self.in_channels = in_channels
36
+ self.out_channels = out_channels
37
+ self.max_num_elements = max_num_elements
38
+
39
+ if mlp is None:
40
+ self.mlp = layers.Dense(out_channels)
41
+ else:
42
+ self.mlp = mlp
43
+
44
+ def build(self, input_shape=None):
45
+ if hasattr(self.mlp, "build"):
46
+ self.mlp.build((None, self.in_channels * self.max_num_elements))
47
+ super().build(input_shape)
48
+
49
+ def reset_parameters(self):
50
+ if hasattr(self.mlp, "reset_parameters"):
51
+ self.mlp.reset_parameters()
52
+
53
+ def call(
54
+ self,
55
+ x,
56
+ index: Optional[any] = None,
57
+ ptr: Optional[any] = None,
58
+ dim_size: Optional[int] = None,
59
+ dim: int = -2,
60
+ **kwargs,
61
+ ):
62
+ x_dense, _ = self.to_dense_batch(
63
+ x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
64
+ max_num_elements=self.max_num_elements,
65
+ )
66
+ B = ops.shape(x_dense)[0]
67
+ flattened = ops.reshape(x_dense, (B, self.max_num_elements * self.in_channels))
68
+ return self.mlp(flattened)
69
+
70
+ def __repr__(self) -> str:
71
+ return (
72
+ f"{self.__class__.__name__}({self.in_channels}, {self.out_channels}, "
73
+ f"max_num_elements={self.max_num_elements})"
74
+ )
75
+
@@ -0,0 +1,154 @@
1
+ import copy
2
+ from typing import Any, Dict, List, Optional, Union
3
+ from keras import layers, ops
4
+
5
+ from .base import Aggregation
6
+
7
+
8
+ class MultiAggregation(Aggregation):
9
+ r"""Performs aggregations with one or more aggregators and combines
10
+ aggregated results, as described in the `"Principal Neighbourhood
11
+ Aggregation for Graph Nets" <https://arxiv.org/abs/2004.05718>`_ and
12
+ `"Adaptive Filters and Aggregator Fusion for Efficient Graph Convolutions"
13
+ <https://arxiv.org/abs/2104.01481>`_ papers.
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import MultiAggregation
19
+
20
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
21
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
22
+
23
+ aggr = MultiAggregation(aggrs=["sum", "mean", "max"])
24
+ out = aggr(x, index=index, dim_size=2)
25
+ print(tuple(out.shape)) # (2, 24)
26
+ ```
27
+ """
28
+
29
+ def __init__(
30
+ self,
31
+ aggrs: List[Union[Aggregation, str]],
32
+ aggrs_kwargs: Optional[List[Dict[str, Any]]] = None,
33
+ mode: Optional[str] = "cat",
34
+ mode_kwargs: Optional[Dict[str, Any]] = None,
35
+ **kwargs,
36
+ ):
37
+ super().__init__(**kwargs)
38
+
39
+ if not isinstance(aggrs, (list, tuple)):
40
+ raise ValueError(f"'aggrs' of '{self.__class__.__name__}' should be a list or tuple.")
41
+
42
+ if len(aggrs) == 0:
43
+ raise ValueError(f"'aggrs' of '{self.__class__.__name__}' should not be empty.")
44
+
45
+ if aggrs_kwargs is None:
46
+ aggrs_kwargs = [{}] * len(aggrs)
47
+ elif len(aggrs) != len(aggrs_kwargs):
48
+ raise ValueError(
49
+ f"'aggrs_kwargs' with invalid length passed to '{self.__class__.__name__}' "
50
+ f"(got '{len(aggrs_kwargs)}', expected '{len(aggrs)}')."
51
+ )
52
+
53
+ from .resolver import aggregation_resolver
54
+ self.aggrs = [
55
+ aggregation_resolver(aggr, **aggr_kw)
56
+ for aggr, aggr_kw in zip(aggrs, aggrs_kwargs)
57
+ ]
58
+
59
+ self.mode = mode
60
+ mode_kwargs = copy.copy(mode_kwargs) or {}
61
+ self.in_channels = mode_kwargs.pop("in_channels", None)
62
+ self.out_channels = mode_kwargs.pop("out_channels", None)
63
+
64
+ if mode in ["proj", "attn"]:
65
+ if len(aggrs) == 1:
66
+ raise ValueError("Multiple aggregations are required for 'proj' or 'attn' combine mode.")
67
+ if self.in_channels is None or self.out_channels is None:
68
+ raise ValueError(f"Combine mode '{mode}' must have `in_channels` and `out_channels` specified.")
69
+
70
+ if isinstance(self.in_channels, int):
71
+ self.in_channels = [self.in_channels] * len(aggrs)
72
+
73
+ if mode == "proj":
74
+ self.lin = layers.Dense(self.out_channels, **mode_kwargs)
75
+ elif mode == "attn":
76
+ from ..dense import HeteroDictLinear
77
+ channels = {str(k): v for k, v in enumerate(self.in_channels)}
78
+ self.lin_heads = HeteroDictLinear(channels, self.out_channels)
79
+ num_heads = mode_kwargs.pop("num_heads", 1)
80
+ self.multihead_attn = layers.MultiHeadAttention(
81
+ num_heads=num_heads,
82
+ key_dim=max(self.out_channels // num_heads, 1),
83
+ **mode_kwargs,
84
+ )
85
+
86
+ def reset_parameters(self):
87
+ for aggr in self.aggrs:
88
+ if hasattr(aggr, "reset_parameters"):
89
+ aggr.reset_parameters()
90
+ if hasattr(self, "lin") and hasattr(self.lin, "reset_parameters"):
91
+ self.lin.reset_parameters()
92
+ if hasattr(self, "lin_heads") and hasattr(self.lin_heads, "reset_parameters"):
93
+ self.lin_heads.reset_parameters()
94
+
95
+ def get_out_channels(self, in_channels: int) -> int:
96
+ if self.out_channels is not None:
97
+ return self.out_channels
98
+ if self.mode == "cat":
99
+ return in_channels * len(self.aggrs)
100
+ return in_channels
101
+
102
+ def call(
103
+ self,
104
+ x,
105
+ index: Optional[any] = None,
106
+ ptr: Optional[any] = None,
107
+ dim_size: Optional[int] = None,
108
+ dim: int = -2,
109
+ **kwargs,
110
+ ):
111
+ outs = [aggr(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim, **kwargs) for aggr in self.aggrs]
112
+ return self.combine(outs)
113
+
114
+ def combine(self, inputs: List[any]):
115
+ if len(inputs) == 1:
116
+ return inputs[0]
117
+
118
+ if self.mode == "cat":
119
+ return ops.concatenate(inputs, axis=-1)
120
+
121
+ if hasattr(self, "lin"):
122
+ return self.lin(ops.concatenate(inputs, axis=-1))
123
+
124
+ if hasattr(self, "multihead_attn"):
125
+ x_dict = {str(k): v for k, v in enumerate(inputs)}
126
+ x_dict = self.lin_heads(x_dict)
127
+ xs = [x_dict[str(key)] for key in range(len(inputs))]
128
+ # xs: [num_aggrs, B, D] -> transpose to [B, num_aggrs, D]
129
+ x_stack = ops.transpose(ops.stack(xs, axis=0), (1, 0, 2))
130
+ attn_out = self.multihead_attn(x_stack, x_stack, x_stack)
131
+ return ops.mean(attn_out, axis=1)
132
+
133
+ stacked = ops.stack(inputs, axis=0) # [num_aggrs, B, D]
134
+ if self.mode == "sum":
135
+ return ops.sum(stacked, axis=0)
136
+ elif self.mode == "mean":
137
+ return ops.mean(stacked, axis=0)
138
+ elif self.mode == "max":
139
+ return ops.max(stacked, axis=0)
140
+ elif self.mode == "min":
141
+ return ops.min(stacked, axis=0)
142
+ elif self.mode == "logsumexp":
143
+ return ops.logsumexp(stacked, axis=0)
144
+ elif self.mode == "std":
145
+ return ops.std(stacked, axis=0)
146
+ elif self.mode == "var":
147
+ return ops.var(stacked, axis=0)
148
+
149
+ raise ValueError(f"Combine mode '{self.mode}' is not supported.")
150
+
151
+ def __repr__(self) -> str:
152
+ aggrs = ",\n".join([f" {aggr}" for aggr in self.aggrs])
153
+ return f"{self.__class__.__name__}([\n{aggrs}\n], mode={self.mode})"
154
+