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
k3_node/io/npz.py ADDED
@@ -0,0 +1,45 @@
1
+ from typing import Any, Dict
2
+ import numpy as np
3
+ import scipy.sparse as sp
4
+ from keras import ops
5
+
6
+ from k3_node.data import Data
7
+ from k3_node.layers.conv.utils import remove_self_loops
8
+ from k3_node.transforms.utils import to_undirected as to_undirected_fn
9
+
10
+
11
+ def read_npz(path: str, to_undirected: bool = True) -> Data:
12
+ """Reads a `.npz` graph archive (e.g. Amazon, Coauthor) and returns a `Data` object."""
13
+ with np.load(path) as f:
14
+ return parse_npz(f, to_undirected=to_undirected)
15
+
16
+
17
+ def parse_npz(f: Dict[str, Any], to_undirected: bool = True) -> Data:
18
+ x_sp = sp.csr_matrix(
19
+ (f["attr_data"], f["attr_indices"], f["attr_indptr"]),
20
+ shape=tuple(f["attr_shape"]),
21
+ )
22
+ x = np.array(x_sp.todense(), dtype=np.float32)
23
+ x[x > 0] = 1.0
24
+
25
+ adj = sp.csr_matrix(
26
+ (f["adj_data"], f["adj_indices"], f["adj_indptr"]),
27
+ shape=tuple(f["adj_shape"]),
28
+ ).tocoo()
29
+
30
+ row = np.array(adj.row, dtype=np.int64)
31
+ col = np.array(adj.col, dtype=np.int64)
32
+ edge_index = np.stack([row, col], axis=0)
33
+
34
+ edge_index, _ = remove_self_loops(edge_index)
35
+ if to_undirected:
36
+ edge_index = to_undirected_fn(edge_index, num_nodes=x.shape[0])
37
+
38
+ y = np.array(f["labels"], dtype=np.int64)
39
+
40
+ return Data(
41
+ x=ops.convert_to_tensor(x, dtype="float32"),
42
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
43
+ y=ops.convert_to_tensor(y, dtype="int64"),
44
+ )
45
+
k3_node/io/off.py ADDED
@@ -0,0 +1,29 @@
1
+ from typing import List
2
+
3
+ import numpy as np
4
+
5
+ from k3_node.data import Data
6
+
7
+
8
+ def parse_off(src: List[str]) -> Data:
9
+ r"""Parses the lines of an OFF (Object File Format) mesh into ``Data(pos, face)``;
10
+ quadrilaterals are split into two triangles."""
11
+ if src[0] == 'OFF':
12
+ src = src[1:]
13
+ else: # some files lack the line break after "OFF"
14
+ src[0] = src[0][3:]
15
+ num_nodes, num_faces = (int(item) for item in src[0].split()[:2])
16
+ pos = np.array([[float(v) for v in line.split()[:3]] for line in src[1:1 + num_nodes]], dtype=np.float32)
17
+ faces = [[int(v) for v in line.strip().split()] for line in src[1 + num_nodes:1 + num_nodes + num_faces]]
18
+ tri = [f[1:4] for f in faces if f[0] == 3]
19
+ for f in faces:
20
+ if f[0] == 4:
21
+ tri += [[f[1], f[2], f[3]], [f[1], f[3], f[4]]]
22
+ face = np.array(tri, dtype=np.int64).reshape(-1, 3).T
23
+ return Data(pos=pos, face=face)
24
+
25
+
26
+ def read_off(path: str) -> Data:
27
+ r"""Reads an OFF (Object File Format) mesh file into ``Data(pos, face)``."""
28
+ with open(path) as f:
29
+ return parse_off(f.read().split('\n')[:-1])
@@ -0,0 +1,98 @@
1
+ import os.path as osp
2
+ import pickle
3
+ import warnings
4
+ from typing import Dict, List, Optional
5
+ import numpy as np
6
+ from keras import ops
7
+
8
+ from k3_node.data import Data
9
+ from k3_node.io.txt_array import read_txt_array
10
+ from k3_node.layers.conv.utils import remove_self_loops
11
+ from k3_node.utils.graph import coalesce
12
+
13
+
14
+ def index_to_mask(index, size: int):
15
+ """Converts 1D index array into a boolean mask."""
16
+ mask = np.zeros(size, dtype=bool)
17
+ idx_np = ops.convert_to_numpy(index)
18
+ mask[idx_np] = True
19
+ return ops.convert_to_tensor(mask, dtype="bool")
20
+
21
+
22
+ def edge_index_from_dict(graph_dict: Dict[int, List[int]], num_nodes: Optional[int] = None):
23
+ rows: List[int] = []
24
+ cols: List[int] = []
25
+ for key, val_list in graph_dict.items():
26
+ for val in val_list:
27
+ rows.append(key)
28
+ cols.append(val)
29
+ if len(rows) == 0:
30
+ return ops.zeros((2, 0), dtype="int64")
31
+
32
+ edge_index = np.array([rows, cols], dtype=np.int64)
33
+ edge_index, _ = remove_self_loops(edge_index)
34
+ edge_index, _ = coalesce(edge_index, num_nodes=num_nodes, sort_by_row=False)
35
+ return ops.convert_to_tensor(edge_index, dtype="int64")
36
+
37
+
38
+ def read_file(folder: str, prefix: str, name: str):
39
+ path = osp.join(folder, f"ind.{prefix.lower()}.{name}")
40
+ if name == "test.index":
41
+ return read_txt_array(path, dtype="int64")
42
+
43
+ with open(path, "rb") as f:
44
+ warnings.filterwarnings("ignore", ".*`scipy.sparse.csr` name.*")
45
+ out = pickle.load(f, encoding="latin1")
46
+
47
+ if name == "graph":
48
+ return out
49
+
50
+ if hasattr(out, "todense"):
51
+ out = out.todense()
52
+ return np.array(out, dtype=np.float32)
53
+
54
+
55
+ def read_planetoid_data(folder: str, prefix: str) -> Data:
56
+ """Reads planetoid citation graph files and returns a `Data` object."""
57
+ names = ["x", "tx", "allx", "y", "ty", "ally", "graph", "test.index"]
58
+ items = [read_file(folder, prefix, name) for name in names]
59
+ x, tx, allx, y, ty, ally, graph, test_index = items
60
+
61
+ test_index_np = ops.convert_to_numpy(test_index).astype(np.int64)
62
+ sorted_test_index = np.sort(test_index_np)
63
+
64
+ train_index = np.arange(y.shape[0], dtype=np.int64)
65
+ val_index = np.arange(y.shape[0], y.shape[0] + 500, dtype=np.int64)
66
+
67
+ if prefix.lower() == "citeseer":
68
+ len_test_indices = int(np.max(test_index_np) - np.min(test_index_np)) + 1
69
+ tx_ext = np.zeros((len_test_indices, tx.shape[1]), dtype=tx.dtype)
70
+ tx_ext[sorted_test_index - np.min(test_index_np), :] = tx
71
+ ty_ext = np.zeros((len_test_indices, ty.shape[1]), dtype=ty.dtype)
72
+ ty_ext[sorted_test_index - np.min(test_index_np), :] = ty
73
+ tx, ty = tx_ext, ty_ext
74
+
75
+ x = np.concatenate([allx, tx], axis=0)
76
+ x[test_index_np] = x[sorted_test_index]
77
+
78
+ y_cat = np.concatenate([ally, ty], axis=0)
79
+ y = np.argmax(y_cat, axis=1).astype(np.int64)
80
+ y[test_index_np] = y[sorted_test_index]
81
+
82
+ num_nodes = y.shape[0]
83
+ train_mask = index_to_mask(train_index, size=num_nodes)
84
+ val_mask = index_to_mask(val_index, size=num_nodes)
85
+ test_mask = index_to_mask(test_index_np, size=num_nodes)
86
+
87
+ edge_index = edge_index_from_dict(graph, num_nodes=num_nodes)
88
+
89
+ data = Data(
90
+ x=ops.convert_to_tensor(x, dtype="float32"),
91
+ edge_index=edge_index,
92
+ y=ops.convert_to_tensor(y, dtype="int64"),
93
+ )
94
+ data.train_mask = train_mask
95
+ data.val_mask = val_mask
96
+ data.test_mask = test_mask
97
+
98
+ return data
k3_node/io/tu.py ADDED
@@ -0,0 +1,137 @@
1
+ import glob
2
+ import os.path as osp
3
+ from typing import Any, Dict, List, Optional, Tuple
4
+ import numpy as np
5
+ from keras import ops
6
+
7
+ from k3_node.data import Data
8
+ from k3_node.io.txt_array import read_txt_array
9
+ from k3_node.layers.conv.utils import remove_self_loops
10
+ from k3_node.utils.graph import coalesce
11
+
12
+
13
+ def np_one_hot(arr: np.ndarray) -> np.ndarray:
14
+ """One-hot encodes a 1D integer array."""
15
+ arr = arr.astype(np.int64)
16
+ if arr.size == 0:
17
+ return np.zeros((0, 0), dtype=np.float32)
18
+ min_val = np.min(arr)
19
+ arr = arr - min_val
20
+ max_val = np.max(arr)
21
+ out = np.zeros((arr.shape[0], max_val + 1), dtype=np.float32)
22
+ out[np.arange(arr.shape[0]), arr] = 1.0
23
+ return out
24
+
25
+
26
+ def read_tu_data(folder: str, prefix: str) -> Tuple[Data, Dict[str, Any], Dict[str, int]]:
27
+ """Reads a TU benchmark dataset from disk and returns `(data, slices, sizes)`."""
28
+ files = sorted(glob.glob(osp.join(folder, f"{prefix}_*.txt")))
29
+ names = [osp.basename(f)[len(prefix) + 1 : -4] for f in files]
30
+
31
+ def read_file(name: str, dtype: str = "float32"):
32
+ path = osp.join(folder, f"{prefix}_{name}.txt")
33
+ return ops.convert_to_numpy(read_txt_array(path, sep=",", dtype=dtype))
34
+
35
+ # Adjacency: 1-indexed in raw TU files
36
+ edge_index_np = read_file("A", dtype="int64")
37
+ if edge_index_np.ndim == 1:
38
+ edge_index_np = edge_index_np.reshape(-1, 2)
39
+ edge_index_np = edge_index_np.T - 1 # 0-indexed (2, E)
40
+
41
+ # Graph indicator: which graph each node belongs to (1-indexed)
42
+ batch_np = read_file("graph_indicator", dtype="int64") - 1 # 0-indexed (N,)
43
+
44
+ node_attribute = np.empty((batch_np.shape[0], 0), dtype=np.float32)
45
+ if "node_attributes" in names:
46
+ node_attribute = read_file("node_attributes", dtype="float32")
47
+ if node_attribute.ndim == 1:
48
+ node_attribute = node_attribute[:, None]
49
+
50
+ node_label = np.empty((batch_np.shape[0], 0), dtype=np.float32)
51
+ if "node_labels" in names:
52
+ node_label_raw = read_file("node_labels", dtype="int64")
53
+ if node_label_raw.ndim == 1:
54
+ node_label_raw = node_label_raw[:, None]
55
+ encoded = [np_one_hot(node_label_raw[:, i]) for i in range(node_label_raw.shape[1])]
56
+ node_label = np.concatenate(encoded, axis=-1)
57
+
58
+ edge_attribute = np.empty((edge_index_np.shape[1], 0), dtype=np.float32)
59
+ if "edge_attributes" in names:
60
+ edge_attribute = read_file("edge_attributes", dtype="float32")
61
+ if edge_attribute.ndim == 1:
62
+ edge_attribute = edge_attribute[:, None]
63
+
64
+ edge_label = np.empty((edge_index_np.shape[1], 0), dtype=np.float32)
65
+ if "edge_labels" in names:
66
+ edge_label_raw = read_file("edge_labels", dtype="int64")
67
+ if edge_label_raw.ndim == 1:
68
+ edge_label_raw = edge_label_raw[:, None]
69
+ encoded = [np_one_hot(edge_label_raw[:, i]) for i in range(edge_label_raw.shape[1])]
70
+ edge_label = np.concatenate(encoded, axis=-1)
71
+
72
+ # Combine attributes and one-hot labels
73
+ x_list = [arr for arr in [node_attribute, node_label] if arr.shape[1] > 0]
74
+ x_np = np.concatenate(x_list, axis=-1) if len(x_list) > 0 else None
75
+
76
+ edge_attr_list = [arr for arr in [edge_attribute, edge_label] if arr.shape[1] > 0]
77
+ edge_attr_np = np.concatenate(edge_attr_list, axis=-1) if len(edge_attr_list) > 0 else None
78
+
79
+ y_np = None
80
+ if "graph_attributes" in names:
81
+ y_np = read_file("graph_attributes", dtype="float32")
82
+ elif "graph_labels" in names:
83
+ y_raw = read_file("graph_labels", dtype="int64")
84
+ _, y_inv = np.unique(y_raw, return_inverse=True)
85
+ y_np = y_inv.astype(np.int64)
86
+
87
+ num_nodes = x_np.shape[0] if x_np is not None else int(np.max(edge_index_np)) + 1
88
+ edge_index_np, edge_attr_np = remove_self_loops(edge_index_np, edge_attr_np)
89
+ edge_index_t, edge_attr_t = coalesce(edge_index_np, edge_attr_np, num_nodes=num_nodes)
90
+ edge_index_np = ops.convert_to_numpy(edge_index_t)
91
+ edge_attr_np = ops.convert_to_numpy(edge_attr_t) if edge_attr_t is not None else None
92
+
93
+ # Convert to tensors
94
+ edge_index = ops.convert_to_tensor(edge_index_np, dtype="int64")
95
+ x = ops.convert_to_tensor(x_np, dtype="float32") if x_np is not None else None
96
+ edge_attr = ops.convert_to_tensor(edge_attr_np, dtype="float32") if edge_attr_np is not None else None
97
+ y = ops.convert_to_tensor(y_np, dtype="float32" if "graph_attributes" in names else "int64") if y_np is not None else None
98
+
99
+ # Compute graph slices
100
+ num_graphs = int(np.max(batch_np)) + 1
101
+ node_counts = np.bincount(batch_np, minlength=num_graphs)
102
+ node_slice = np.pad(np.cumsum(node_counts), (1, 0))
103
+
104
+ row = edge_index_np[0]
105
+ edge_batch = batch_np[row]
106
+ edge_counts = np.bincount(edge_batch, minlength=num_graphs)
107
+ edge_slice = np.pad(np.cumsum(edge_counts), (1, 0))
108
+
109
+ # Shift edge indices so each graph starts at 0
110
+ shift = node_slice[edge_batch]
111
+ edge_index_shifted = edge_index_np - shift[None, :]
112
+ edge_index = ops.convert_to_tensor(edge_index_shifted, dtype="int64")
113
+
114
+ data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr, y=y)
115
+ slices: Dict[str, Any] = {
116
+ "edge_index": ops.convert_to_tensor(edge_slice, dtype="int64"),
117
+ }
118
+ if x is not None:
119
+ slices["x"] = ops.convert_to_tensor(node_slice, dtype="int64")
120
+ else:
121
+ data.num_nodes = int(batch_np.shape[0])
122
+ if edge_attr is not None:
123
+ slices["edge_attr"] = ops.convert_to_tensor(edge_slice, dtype="int64")
124
+ if y is not None:
125
+ if y_np.shape[0] == batch_np.shape[0]:
126
+ slices["y"] = ops.convert_to_tensor(node_slice, dtype="int64")
127
+ else:
128
+ slices["y"] = ops.convert_to_tensor(np.arange(num_graphs + 1, dtype=np.int64), dtype="int64")
129
+
130
+ sizes = {
131
+ "num_node_attributes": node_attribute.shape[-1],
132
+ "num_node_labels": node_label.shape[-1],
133
+ "num_edge_attributes": edge_attribute.shape[-1],
134
+ "num_edge_labels": edge_label.shape[-1],
135
+ }
136
+
137
+ return data, slices, sizes
@@ -0,0 +1,58 @@
1
+ from typing import List, Optional, Union
2
+ import numpy as np
3
+ from keras import ops
4
+
5
+
6
+ def parse_txt_array(
7
+ src: List[str],
8
+ sep: Optional[str] = None,
9
+ start: int = 0,
10
+ end: Optional[int] = None,
11
+ dtype: Optional[str] = None,
12
+ ):
13
+ """Parses a list of string rows into a tensor."""
14
+ lines = [line.strip() for line in src if line.strip()]
15
+ if len(lines) == 0:
16
+ return ops.zeros((0,), dtype=dtype or "float32")
17
+
18
+ split_lines = []
19
+ is_float = False
20
+ for line in lines:
21
+ parts = line.split(sep)[start:end]
22
+ row = []
23
+ for x in parts:
24
+ x_str = x.strip()
25
+ if not x_str:
26
+ continue
27
+ if "." in x_str or "e" in x_str.lower():
28
+ is_float = True
29
+ row.append(float(x_str))
30
+ else:
31
+ try:
32
+ row.append(int(x_str))
33
+ except ValueError:
34
+ is_float = True
35
+ row.append(float(x_str))
36
+ split_lines.append(row)
37
+
38
+ if dtype is None:
39
+ dtype = "float32" if is_float else "int64"
40
+
41
+ arr = np.array(split_lines, dtype=dtype)
42
+ if arr.ndim > 1 and arr.shape[1] == 1:
43
+ arr = np.squeeze(arr, axis=1)
44
+ return ops.convert_to_tensor(arr, dtype=dtype)
45
+
46
+
47
+ def read_txt_array(
48
+ path: str,
49
+ sep: Optional[str] = None,
50
+ start: int = 0,
51
+ end: Optional[int] = None,
52
+ dtype: Optional[str] = None,
53
+ ):
54
+ """Reads a text array from a file and returns a tensor."""
55
+ with open(path, "r", encoding="utf-8") as f:
56
+ src = f.read().split("\n")
57
+ return parse_txt_array(src, sep, start, end, dtype)
58
+
@@ -0,0 +1,14 @@
1
+ """
2
+ `k3_node.layers` module provides access to various layers for building
3
+ graph neural networks.
4
+ """
5
+
6
+ from .conv import *
7
+ from .norm import *
8
+ from .aggr import *
9
+ from .attention import *
10
+ from .dense import *
11
+ from .pool import *
12
+ from .unpool import *
13
+ from .kge import *
14
+ from .functional import *
@@ -0,0 +1,70 @@
1
+ r"""Aggregation operators for graph neural networks."""
2
+
3
+ from .base import Aggregation, from_dense_batch, ptr2index, to_dense_adj, to_dense_batch
4
+ from .basic import (
5
+ MaxAggregation,
6
+ MeanAggregation,
7
+ MinAggregation,
8
+ MulAggregation,
9
+ PowerMeanAggregation,
10
+ SoftmaxAggregation,
11
+ StdAggregation,
12
+ SumAggregation,
13
+ VarAggregation,
14
+ )
15
+ from .quantile import MedianAggregation, QuantileAggregation
16
+ from .attention import AttentionalAggregation
17
+ from .set2set import Set2Set
18
+ from .scaler import DegreeScalerAggregation
19
+ from .sort import SortAggregation
20
+ from .multi import MultiAggregation
21
+ from .deep_sets import DeepSetsAggregation
22
+ from .mlp import MLPAggregation
23
+ from .lstm import LSTMAggregation
24
+ from .gru import GRUAggregation
25
+ from .set_transformer import SetTransformerAggregation
26
+ from .gmt import GraphMultisetTransformer
27
+ from .variance_preserving import VariancePreservingAggregation
28
+ from .patch_transformer import PatchTransformerAggregation
29
+ from .lcm import LCMAggregation
30
+ from .equilibrium import EquilibriumAggregation
31
+ from .fused import FusedAggregation
32
+ from .resolver import aggregation_resolver
33
+
34
+ __all__ = [
35
+ "Aggregation",
36
+ "ptr2index",
37
+ "to_dense_batch",
38
+ "from_dense_batch",
39
+ "to_dense_adj",
40
+ "SumAggregation",
41
+ "MeanAggregation",
42
+ "MaxAggregation",
43
+ "MinAggregation",
44
+ "MulAggregation",
45
+ "VarAggregation",
46
+ "StdAggregation",
47
+ "SoftmaxAggregation",
48
+ "PowerMeanAggregation",
49
+ "QuantileAggregation",
50
+ "MedianAggregation",
51
+ "AttentionalAggregation",
52
+ "Set2Set",
53
+ "DegreeScalerAggregation",
54
+ "SortAggregation",
55
+ "MultiAggregation",
56
+ "DeepSetsAggregation",
57
+ "MLPAggregation",
58
+ "LSTMAggregation",
59
+ "GRUAggregation",
60
+ "SetTransformerAggregation",
61
+ "GraphMultisetTransformer",
62
+ "VariancePreservingAggregation",
63
+ "PatchTransformerAggregation",
64
+ "LCMAggregation",
65
+ "EquilibriumAggregation",
66
+ "FusedAggregation",
67
+ "aggregation_resolver",
68
+ ]
69
+
70
+ classes = __all__
@@ -0,0 +1,77 @@
1
+ from typing import Optional
2
+ from keras import ops
3
+
4
+ from .base import Aggregation
5
+ from k3_node.ops.segment import segment_max, segment_sum
6
+
7
+
8
+ class AttentionalAggregation(Aggregation):
9
+ r"""The soft attention aggregation layer from the `"Graph Matching Networks
10
+ for Learning the Similarity of Graph Structured Objects"
11
+ <https://arxiv.org/abs/1904.12787>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ import keras
17
+ from k3_node.layers import AttentionalAggregation
18
+
19
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
20
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
21
+
22
+ aggr = AttentionalAggregation(gate_nn=keras.layers.Dense(1), nn=keras.layers.Dense(16))
23
+ out = aggr(x, index=index, dim_size=2) # attention-weighted sum per set
24
+ print(tuple(out.shape)) # (2, 16)
25
+ ```
26
+ """
27
+
28
+ def __init__(
29
+ self,
30
+ gate_nn,
31
+ nn: Optional[any] = None,
32
+ **kwargs,
33
+ ):
34
+ super().__init__(**kwargs)
35
+ self.gate_nn = gate_nn
36
+ self.nn = nn
37
+
38
+ def reset_parameters(self):
39
+ if hasattr(self.gate_nn, "reset_parameters"):
40
+ self.gate_nn.reset_parameters()
41
+ if self.nn is not None and hasattr(self.nn, "reset_parameters"):
42
+ self.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
+ gate = self.gate_nn(x)
54
+ if self.nn is not None:
55
+ x = self.nn(x)
56
+
57
+ if ptr is not None and index is None:
58
+ from .base import ptr2index
59
+ index = ptr2index(ptr)
60
+
61
+ index = ops.cast(index, dtype="int32")
62
+ if dim_size is None: # a tensor while tracing; don't test its truth value
63
+ dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
64
+
65
+ # Graph-wise softmax over groups
66
+ max_val = segment_max(gate, index, num_segments=dim_size)
67
+ max_exp = ops.take(max_val, index, axis=0)
68
+ exp_gate = ops.exp(gate - max_exp)
69
+ sum_exp = segment_sum(exp_gate, index, num_segments=dim_size)
70
+ sum_exp_exp = ops.take(sum_exp, index, axis=0)
71
+ alpha = exp_gate / ops.maximum(sum_exp_exp, 1e-12)
72
+
73
+ return self.reduce(alpha * x, index, ptr, dim_size, dim, reduce="sum")
74
+
75
+ def __repr__(self) -> str:
76
+ return f"{self.__class__.__name__}(gate_nn={self.gate_nn}, nn={self.nn})"
77
+