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,137 @@
1
+ import math
2
+ from typing import List, Optional, Union
3
+ from keras import layers, ops
4
+ import numpy as np
5
+
6
+ from .base import Aggregation
7
+ from .utils import MultiheadAttentionBlock
8
+
9
+
10
+ class PatchTransformerAggregation(Aggregation):
11
+ r"""Performs patch transformer aggregation in which the elements to
12
+ aggregate are processed by multi-head attention blocks across patches.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import PatchTransformerAggregation
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 = PatchTransformerAggregation(in_channels=8, out_channels=16, patch_size=2, hidden_channels=8)
23
+ out = aggr(x, index=index, dim_size=2)
24
+ print(tuple(out.shape)) # (2, 16)
25
+ ```
26
+ """
27
+
28
+ def __init__(
29
+ self,
30
+ in_channels: int,
31
+ out_channels: int,
32
+ patch_size: int,
33
+ hidden_channels: int,
34
+ num_transformer_blocks: int = 1,
35
+ heads: int = 1,
36
+ dropout: float = 0.0,
37
+ aggr: Union[str, List[str]] = "mean",
38
+ **kwargs,
39
+ ):
40
+ super().__init__(**kwargs)
41
+ self.in_channels = in_channels
42
+ self.out_channels = out_channels
43
+ self.patch_size = patch_size
44
+ self.hidden_channels = hidden_channels
45
+ self.aggrs = [aggr] if isinstance(aggr, str) else list(aggr)
46
+
47
+ self.lin = layers.Dense(hidden_channels)
48
+ self.pad_projector = layers.Dense(hidden_channels)
49
+ self.blocks = [
50
+ MultiheadAttentionBlock(
51
+ channels=hidden_channels,
52
+ heads=heads,
53
+ layer_norm=True,
54
+ dropout=dropout,
55
+ )
56
+ for _ in range(num_transformer_blocks)
57
+ ]
58
+ self.fc = layers.Dense(out_channels)
59
+
60
+ def reset_parameters(self):
61
+ self.lin.reset_parameters()
62
+ self.pad_projector.reset_parameters()
63
+ for block in self.blocks:
64
+ block.reset_parameters()
65
+ self.fc.reset_parameters()
66
+
67
+ def call(
68
+ self,
69
+ x,
70
+ index: Optional[any] = None,
71
+ ptr: Optional[any] = None,
72
+ dim_size: Optional[int] = None,
73
+ dim: int = -2,
74
+ max_num_elements: Optional[int] = None,
75
+ training: bool = False,
76
+ **kwargs,
77
+ ):
78
+ if max_num_elements is None:
79
+ from k3_node.layers.conv.utils import is_tracing
80
+ if is_tracing(x) or is_tracing(index):
81
+ if hasattr(x, "shape") and x.shape[0] is not None:
82
+ max_num_elements = int(x.shape[0])
83
+ else:
84
+ max_num_elements = 16
85
+ elif ptr is not None:
86
+ ptr_np = ops.convert_to_numpy(ptr)
87
+ count = ptr_np[1:] - ptr_np[:-1]
88
+ max_num_elements = int(np.max(count)) if len(count) > 0 else 1
89
+ else:
90
+ idx_np = ops.convert_to_numpy(index).astype(np.int64)
91
+ counts = np.bincount(idx_np)
92
+ max_num_elements = int(np.max(counts)) if len(counts) > 0 else 1
93
+
94
+ # Ensure max_num_elements is a multiple of patch_size
95
+ num_patches = max(math.ceil(max_num_elements / self.patch_size), 1)
96
+ target_elements = num_patches * self.patch_size
97
+
98
+ x_dense, _ = self.to_dense_batch(
99
+ x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
100
+ max_num_elements=target_elements,
101
+ )
102
+
103
+ B = ops.shape(x_dense)[0]
104
+ x_proj = self.lin(x_dense) # [B, target_elements, hidden_channels]
105
+
106
+ # Reshape to patches: [B, num_patches, patch_size * hidden_channels]
107
+ x_patches = ops.reshape(x_proj, (B, num_patches, self.patch_size * self.hidden_channels))
108
+ x_patches = self.pad_projector(x_patches) # [B, num_patches, hidden_channels]
109
+
110
+ # Process through transformer blocks
111
+ for block in self.blocks:
112
+ x_patches = block(x_patches, x_patches, training=training)
113
+
114
+ outs = []
115
+ for aggr_mode in self.aggrs:
116
+ if aggr_mode == "mean":
117
+ outs.append(ops.mean(x_patches, axis=1))
118
+ elif aggr_mode == "sum":
119
+ outs.append(ops.sum(x_patches, axis=1))
120
+ elif aggr_mode == "max":
121
+ outs.append(ops.max(x_patches, axis=1))
122
+ elif aggr_mode == "min":
123
+ outs.append(ops.min(x_patches, axis=1))
124
+ elif aggr_mode == "var":
125
+ outs.append(ops.var(x_patches, axis=1))
126
+ elif aggr_mode == "std":
127
+ outs.append(ops.std(x_patches, axis=1))
128
+
129
+ combined = ops.concatenate(outs, axis=-1) if len(outs) > 1 else outs[0]
130
+ return self.fc(combined)
131
+
132
+ def __repr__(self) -> str:
133
+ return (
134
+ f"{self.__class__.__name__}({self.in_channels}, "
135
+ f"{self.out_channels}, patch_size={self.patch_size})"
136
+ )
137
+
@@ -0,0 +1,125 @@
1
+ from typing import List, Optional, Union
2
+ from keras import ops
3
+ import numpy as np
4
+
5
+ from .base import Aggregation
6
+
7
+
8
+ class QuantileAggregation(Aggregation):
9
+ r"""An aggregation operator that returns the feature-wise :math:`q`-th
10
+ quantile of a set :math:`\mathcal{X}`.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import QuantileAggregation
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 = QuantileAggregation(q=0.75)
21
+ out = aggr(x, index=index, dim_size=2)
22
+ print(tuple(out.shape)) # (2, 8)
23
+ ```
24
+ """
25
+ interpolations = {"linear", "lower", "higher", "nearest", "midpoint"}
26
+
27
+ def __init__(
28
+ self,
29
+ q: Union[float, List[float]],
30
+ interpolation: str = "linear",
31
+ fill_value: float = 0.0,
32
+ **kwargs,
33
+ ):
34
+ super().__init__(**kwargs)
35
+
36
+ qs = [q] if not isinstance(q, (list, tuple)) else list(q)
37
+ if len(qs) == 0:
38
+ raise ValueError("Provide at least one quantile value for `q`.")
39
+ if not all(0.0 <= quantile <= 1.0 for quantile in qs):
40
+ raise ValueError("`q` must be in the range [0, 1].")
41
+ if interpolation not in self.interpolations:
42
+ raise ValueError(f"Invalid interpolation method got ('{interpolation}')")
43
+
44
+ self.q = qs
45
+ self.interpolation = interpolation
46
+ self.fill_value = fill_value
47
+
48
+ def call(
49
+ self,
50
+ x,
51
+ index: Optional[any] = None,
52
+ ptr: Optional[any] = None,
53
+ dim_size: Optional[int] = None,
54
+ dim: int = -2,
55
+ **kwargs,
56
+ ):
57
+ self.assert_index_present(index)
58
+ from k3_node.ops.segment import segment_sum
59
+
60
+ # Sort every set's values in a dense [sets, max_size, features] tensor; the padding sorts
61
+ # last (+inf) and is then zeroed. Static shapes and differentiable, like PyG's version.
62
+ dense, _ = self.to_dense_batch(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
63
+ fill_value=float("inf"))
64
+ dense = ops.sort(dense, axis=1)
65
+ dense = ops.where(ops.isinf(dense), ops.zeros_like(dense), dense)
66
+
67
+ index_i = ops.cast(index, "int32")
68
+ count = segment_sum(ops.ones_like(index_i), index_i, num_segments=ops.shape(dense)[0])
69
+ count_f = ops.cast(count, dense.dtype)
70
+ last = ops.maximum(count - 1, 0)
71
+
72
+ def gather(position): # the value at `position` of every set, per feature
73
+ position = ops.minimum(ops.maximum(ops.cast(position, "int32"), 0), last)
74
+ position = ops.broadcast_to(ops.reshape(position, (-1, 1, 1)),
75
+ (ops.shape(dense)[0], 1, ops.shape(dense)[2]))
76
+ return ops.take_along_axis(dense, position, axis=1)[:, 0]
77
+
78
+ outs = []
79
+ for q_val in self.q:
80
+ q_point = q_val * (count_f - 1.0)
81
+ if self.interpolation == "lower":
82
+ quantile = gather(ops.floor(q_point))
83
+ elif self.interpolation == "higher":
84
+ quantile = gather(ops.ceil(q_point))
85
+ elif self.interpolation == "nearest":
86
+ quantile = gather(ops.round(q_point))
87
+ else:
88
+ low, high = gather(ops.floor(q_point)), gather(ops.ceil(q_point))
89
+ if self.interpolation == "linear":
90
+ frac = ops.expand_dims(q_point - ops.floor(q_point), -1)
91
+ quantile = low + (high - low) * frac
92
+ else: # midpoint
93
+ quantile = 0.5 * low + 0.5 * high
94
+ empty = ops.expand_dims(count == 0, -1)
95
+ outs.append(ops.where(empty, ops.cast(self.fill_value, quantile.dtype), quantile))
96
+ return outs[0] if len(outs) == 1 else ops.concatenate(outs, axis=-1)
97
+
98
+ def __repr__(self) -> str:
99
+ q_str = self.q[0] if len(self.q) == 1 else self.q
100
+ return f"{self.__class__.__name__}(q={q_str})"
101
+
102
+
103
+ class MedianAggregation(QuantileAggregation):
104
+ r"""An aggregation operator that returns the feature-wise median of a set.
105
+
106
+ Example:
107
+ ```python
108
+ import numpy as np
109
+ from k3_node.layers import MedianAggregation
110
+
111
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
112
+ index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
113
+
114
+ aggr = MedianAggregation()
115
+ out = aggr(x, index=index, dim_size=2)
116
+ print(tuple(out.shape)) # (2, 8)
117
+ ```
118
+ """
119
+
120
+ def __init__(self, fill_value: float = 0.0, **kwargs):
121
+ super().__init__(0.5, "lower", fill_value=fill_value, **kwargs)
122
+
123
+ def __repr__(self) -> str:
124
+ return f"{self.__class__.__name__}()"
125
+
@@ -0,0 +1,68 @@
1
+ from typing import Union
2
+
3
+ from .base import Aggregation
4
+
5
+
6
+ def aggregation_resolver(
7
+ query: Union[str, Aggregation],
8
+ *args,
9
+ **kwargs,
10
+ ) -> Aggregation:
11
+ r"""Resolves an aggregation string or instance to an `Aggregation` object.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import aggregation_resolver
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 = aggregation_resolver("mean") # build an Aggregation from its name
22
+ print(type(aggr).__name__, tuple(aggr(x, index=index, dim_size=2).shape)) # MeanAggregation (2, 8)
23
+ ```
24
+ """
25
+ if isinstance(query, Aggregation):
26
+ return query
27
+
28
+ if not isinstance(query, str):
29
+ raise ValueError(f"Expected string or Aggregation instance, got {type(query)}")
30
+
31
+ query_norm = query.lower().strip()
32
+
33
+ from .basic import (
34
+ MaxAggregation,
35
+ MeanAggregation,
36
+ MinAggregation,
37
+ MulAggregation,
38
+ PowerMeanAggregation,
39
+ SoftmaxAggregation,
40
+ StdAggregation,
41
+ SumAggregation,
42
+ VarAggregation,
43
+ )
44
+ from .quantile import MedianAggregation, QuantileAggregation
45
+ from .variance_preserving import VariancePreservingAggregation
46
+
47
+ AGGR_DICT = {
48
+ "sum": SumAggregation,
49
+ "add": SumAggregation,
50
+ "mean": MeanAggregation,
51
+ "max": MaxAggregation,
52
+ "min": MinAggregation,
53
+ "mul": MulAggregation,
54
+ "var": VarAggregation,
55
+ "std": StdAggregation,
56
+ "softmax": SoftmaxAggregation,
57
+ "powermean": PowerMeanAggregation,
58
+ "median": MedianAggregation,
59
+ "quantile": QuantileAggregation,
60
+ "variance_preserving": VariancePreservingAggregation,
61
+ "vpa": VariancePreservingAggregation,
62
+ }
63
+
64
+ if query_norm in AGGR_DICT:
65
+ return AGGR_DICT[query_norm](*args, **kwargs)
66
+
67
+ raise ValueError(f"Could not resolve aggregation '{query}'")
68
+
@@ -0,0 +1,133 @@
1
+ from typing import Any, Dict, List, Optional, Union
2
+ from keras import initializers, ops
3
+ import numpy as np
4
+
5
+ from .base import Aggregation
6
+ from k3_node.ops.segment import segment_sum
7
+ from k3_node.ops.creation import full
8
+
9
+
10
+ class DegreeScalerAggregation(Aggregation):
11
+ r"""Combines one or more aggregators and transforms its output with one or
12
+ more scalers as introduced in the `"Principal Neighbourhood Aggregation for
13
+ Graph Nets" <https://arxiv.org/abs/2004.05718>`_ paper.
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import DegreeScalerAggregation
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
+ deg = np.array([0, 3, 5, 2]) # in-degree histogram of the training graphs
24
+ aggr = DegreeScalerAggregation(aggr=["mean", "max"], scaler=["identity", "amplification"], deg=deg)
25
+ out = aggr(x, index=index, dim_size=2)
26
+ print(tuple(out.shape)) # (2, 32): 2 aggregators x 2 scalers x 8 features
27
+ ```
28
+ """
29
+
30
+ def __init__(
31
+ self,
32
+ aggr: Union[str, List[str], Aggregation],
33
+ scaler: Union[str, List[str]],
34
+ deg,
35
+ train_norm: bool = False,
36
+ aggr_kwargs: Optional[List[Dict[str, Any]]] = None,
37
+ **kwargs,
38
+ ):
39
+ super().__init__(**kwargs)
40
+
41
+ from .resolver import aggregation_resolver
42
+ from .multi import MultiAggregation
43
+
44
+ if isinstance(aggr, (str, Aggregation)):
45
+ self.aggr = aggregation_resolver(aggr, **(aggr_kwargs or {}))
46
+ elif isinstance(aggr, (tuple, list)):
47
+ self.aggr = MultiAggregation(aggr, aggr_kwargs)
48
+ else:
49
+ raise ValueError(
50
+ f"Only strings, list, tuples and instances of "
51
+ f"`Aggregation` are valid aggregation schemes (got '{type(aggr)}')"
52
+ )
53
+
54
+ self.scaler = [scaler] if isinstance(scaler, str) else list(scaler)
55
+
56
+ deg_np = ops.convert_to_numpy(deg).astype(np.float32)
57
+ N = float(np.sum(deg_np))
58
+ bin_degree = np.arange(len(deg_np), dtype=np.float32)
59
+
60
+ self.init_avg_deg_lin = float(np.sum(bin_degree * deg_np)) / max(N, 1.0)
61
+ self.init_avg_deg_log = float(np.sum(np.log(bin_degree + 1.0) * deg_np)) / max(N, 1.0)
62
+ self.train_norm = train_norm
63
+
64
+ if train_norm:
65
+ self.avg_deg_lin = self.add_weight(
66
+ shape=(1,),
67
+ initializer=initializers.Constant(self.init_avg_deg_lin),
68
+ trainable=True,
69
+ name="avg_deg_lin",
70
+ )
71
+ self.avg_deg_log = self.add_weight(
72
+ shape=(1,),
73
+ initializer=initializers.Constant(self.init_avg_deg_log),
74
+ trainable=True,
75
+ name="avg_deg_log",
76
+ )
77
+ else:
78
+ self.avg_deg_lin = self.init_avg_deg_lin
79
+ self.avg_deg_log = self.init_avg_deg_log
80
+
81
+ def reset_parameters(self):
82
+ if hasattr(self.aggr, "reset_parameters"):
83
+ self.aggr.reset_parameters()
84
+ if self.train_norm:
85
+ self.avg_deg_lin.assign(full((1,), self.init_avg_deg_lin, dtype=self.avg_deg_lin.dtype))
86
+ self.avg_deg_log.assign(full((1,), self.init_avg_deg_log, dtype=self.avg_deg_log.dtype))
87
+
88
+ def call(
89
+ self,
90
+ x,
91
+ index: Optional[any] = None,
92
+ ptr: Optional[any] = None,
93
+ dim_size: Optional[int] = None,
94
+ dim: int = -2,
95
+ **kwargs,
96
+ ):
97
+ self.assert_index_present(index)
98
+
99
+ out = self.aggr(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim)
100
+
101
+ index = ops.cast(index, dtype="int32")
102
+ if dim_size is None: # a tensor while tracing; don't test its truth value
103
+ dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
104
+
105
+ # Compute degree per index
106
+ ones = ops.ones((ops.shape(index)[0],), dtype=out.dtype)
107
+ deg = segment_sum(ones, index, num_segments=dim_size)
108
+ deg = ops.reshape(deg, (dim_size,) + (1,) * (len(ops.shape(out)) - 1))
109
+
110
+ avg_deg_log = self.avg_deg_log
111
+ avg_deg_lin = self.avg_deg_lin
112
+
113
+ outs = []
114
+ for scaler in self.scaler:
115
+ if scaler == "identity":
116
+ out_scaler = out
117
+ elif scaler == "amplification":
118
+ out_scaler = out * (ops.log(deg + 1.0) / avg_deg_log)
119
+ elif scaler == "attenuation":
120
+ out_scaler = out * (avg_deg_log / ops.log(ops.maximum(deg, 1.0) + 1.0))
121
+ elif scaler == "linear":
122
+ out_scaler = out * (deg / avg_deg_lin)
123
+ elif scaler == "inverse_linear":
124
+ out_scaler = out * (avg_deg_lin / ops.maximum(deg, 1.0))
125
+ else:
126
+ raise ValueError(f"Unknown scaler '{scaler}'")
127
+ outs.append(out_scaler)
128
+
129
+ return ops.concatenate(outs, axis=-1) if len(outs) > 1 else outs[0]
130
+
131
+ def __repr__(self) -> str:
132
+ return f"{self.__class__.__name__}(aggr={self.aggr}, scaler={self.scaler})"
133
+
@@ -0,0 +1,87 @@
1
+ from typing import Optional
2
+ from keras import layers, ops
3
+
4
+ from .base import Aggregation
5
+ from k3_node.ops.segment import segment_max, segment_sum
6
+
7
+
8
+ class Set2Set(Aggregation):
9
+ r"""The Set2Set aggregation operator based on iterative content-based
10
+ attention, as described in the `"Order Matters: Sequence to sequence for
11
+ Sets" <https://arxiv.org/abs/1511.06391>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import Set2Set
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 = Set2Set(in_channels=8, processing_steps=2)
22
+ out = aggr(x, index=index, dim_size=2)
23
+ print(tuple(out.shape)) # (2, 16)
24
+ ```
25
+ """
26
+
27
+ def __init__(self, in_channels: int, processing_steps: int, **kwargs):
28
+ super().__init__(**kwargs)
29
+ self.in_channels = in_channels
30
+ self.out_channels = 2 * in_channels
31
+ self.processing_steps = processing_steps
32
+ self.lstm_cell = layers.LSTMCell(in_channels)
33
+
34
+ def build(self, input_shape=None):
35
+ self.lstm_cell.build((None, self.out_channels))
36
+ super().build(input_shape)
37
+
38
+ def reset_parameters(self):
39
+ self.lstm_cell.reset_parameters()
40
+
41
+ def call(
42
+ self,
43
+ x,
44
+ index: Optional[any] = None,
45
+ ptr: Optional[any] = None,
46
+ dim_size: Optional[int] = None,
47
+ dim: int = -2,
48
+ **kwargs,
49
+ ):
50
+ if ptr is not None and index is None:
51
+ from .base import ptr2index
52
+ index = ptr2index(ptr)
53
+
54
+ self.assert_index_present(index)
55
+ self.assert_two_dimensional_input(x, dim)
56
+
57
+ index = ops.cast(index, dtype="int32")
58
+ if dim_size is None: # a tensor while tracing; don't test its truth value
59
+ dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
60
+
61
+ # Initial hidden states: [dim_size, in_channels]
62
+ h = [
63
+ ops.zeros((dim_size, self.in_channels), dtype=x.dtype),
64
+ ops.zeros((dim_size, self.in_channels), dtype=x.dtype),
65
+ ]
66
+ q_star = ops.zeros((dim_size, self.out_channels), dtype=x.dtype)
67
+
68
+ for _ in range(self.processing_steps):
69
+ q, h = self.lstm_cell(q_star, h)
70
+ q_taken = ops.take(q, index, axis=0)
71
+ e = ops.sum(x * q_taken, axis=-1, keepdims=True)
72
+
73
+ max_e = segment_max(e, index, num_segments=dim_size)
74
+ max_e_exp = ops.take(max_e, index, axis=0)
75
+ exp_e = ops.exp(e - max_e_exp)
76
+ sum_exp = segment_sum(exp_e, index, num_segments=dim_size)
77
+ sum_exp_exp = ops.take(sum_exp, index, axis=0)
78
+ a = exp_e / ops.maximum(sum_exp_exp, 1e-12)
79
+
80
+ r = segment_sum(a * x, index, num_segments=dim_size)
81
+ q_star = ops.concatenate([q, r], axis=-1)
82
+
83
+ return q_star
84
+
85
+ def __repr__(self) -> str:
86
+ return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels})"
87
+
@@ -0,0 +1,107 @@
1
+ from typing import Optional
2
+ from keras import layers, ops
3
+
4
+ from .base import Aggregation
5
+ from .utils import PoolingByMultiheadAttention, SetAttentionBlock
6
+
7
+
8
+ class SetTransformerAggregation(Aggregation):
9
+ r"""Performs "Set Transformer" aggregation in which the elements to
10
+ aggregate are processed by multi-head attention blocks, as described in
11
+ the `"Graph Neural Networks with Adaptive Readouts"
12
+ <https://arxiv.org/abs/2211.04952>`_ paper.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import SetTransformerAggregation
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 = SetTransformerAggregation(channels=8, num_seed_points=2)
23
+ out = aggr(x, index=index, dim_size=2)
24
+ print(tuple(out.shape)) # (2, 16)
25
+ ```
26
+ """
27
+
28
+ def __init__(
29
+ self,
30
+ channels: int,
31
+ num_seed_points: int = 1,
32
+ num_encoder_blocks: int = 1,
33
+ num_decoder_blocks: int = 1,
34
+ heads: int = 1,
35
+ concat: bool = True,
36
+ layer_norm: bool = False,
37
+ dropout: float = 0.0,
38
+ **kwargs,
39
+ ):
40
+ super().__init__(**kwargs)
41
+ self.channels = channels
42
+ self.num_seed_points = num_seed_points
43
+ self.heads = heads
44
+ self.concat = concat
45
+ self.layer_norm = layer_norm
46
+ self.dropout = dropout
47
+
48
+ self.encoders = [
49
+ SetAttentionBlock(channels, heads, layer_norm, dropout)
50
+ for _ in range(num_encoder_blocks)
51
+ ]
52
+ self.pma = PoolingByMultiheadAttention(
53
+ channels, num_seed_points, heads, layer_norm, dropout
54
+ )
55
+ self.decoders = [
56
+ SetAttentionBlock(channels, heads, layer_norm, dropout)
57
+ for _ in range(num_decoder_blocks)
58
+ ]
59
+
60
+ def reset_parameters(self):
61
+ for encoder in self.encoders:
62
+ encoder.reset_parameters()
63
+ self.pma.reset_parameters()
64
+ for decoder in self.decoders:
65
+ decoder.reset_parameters()
66
+
67
+ def call(
68
+ self,
69
+ x,
70
+ index: Optional[any] = None,
71
+ ptr: Optional[any] = None,
72
+ dim_size: Optional[int] = None,
73
+ dim: int = -2,
74
+ max_num_elements: Optional[int] = None,
75
+ training: bool = False,
76
+ **kwargs,
77
+ ):
78
+ x_dense, mask = self.to_dense_batch(
79
+ x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
80
+ max_num_elements=max_num_elements,
81
+ )
82
+
83
+ for encoder in self.encoders:
84
+ x_dense = encoder(x_dense, mask=mask, training=training)
85
+
86
+ x_dense = self.pma(x_dense, mask=mask, training=training)
87
+
88
+ for decoder in self.decoders:
89
+ x_dense = decoder(x_dense, training=training)
90
+
91
+ # Handle NaNs if any
92
+ x_dense = ops.where(ops.isnan(x_dense), 0.0, x_dense)
93
+
94
+ if self.concat:
95
+ B = ops.shape(x_dense)[0]
96
+ return ops.reshape(x_dense, (B, self.num_seed_points * self.channels))
97
+ else:
98
+ return ops.mean(x_dense, axis=1)
99
+
100
+ def __repr__(self) -> str:
101
+ return (
102
+ f"{self.__class__.__name__}({self.channels}, "
103
+ f"num_seed_points={self.num_seed_points}, "
104
+ f"heads={self.heads}, layer_norm={self.layer_norm}, "
105
+ f"dropout={self.dropout})"
106
+ )
107
+