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,119 @@
1
+ from typing import Optional, Tuple
2
+ from keras import ops
3
+ import numpy as np
4
+ from k3_node.ops.segment import segment_max, segment_sum
5
+ from k3_node.ops.host import to_numpy
6
+
7
+
8
+ def pool_edge(
9
+ cluster,
10
+ edge_index,
11
+ edge_attr: Optional[any] = None,
12
+ reduce: str = "sum",
13
+ ) -> Tuple[any, Optional[any]]:
14
+ r"""Pools edge indices and attributes based on cluster assignments.
15
+
16
+ Example:
17
+ ```python
18
+ import numpy as np
19
+ from k3_node.layers import pool_edge
20
+
21
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
22
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
23
+ cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
24
+
25
+ edge_attr = np.random.rand(30, 3).astype("float32")
26
+ edge_index_pool, edge_attr_pool = pool_edge(cluster, edge_index, edge_attr) # coarsened, deduplicated edges
27
+ print(edge_index_pool.shape[0], edge_attr_pool.shape[1]) # 2 3
28
+ ```
29
+ """
30
+ cluster_np = to_numpy(cluster)
31
+ edge_index_np = to_numpy(edge_index)
32
+
33
+ row = cluster_np[edge_index_np[0]]
34
+ col = cluster_np[edge_index_np[1]]
35
+
36
+ # Remove self-loops
37
+ non_loop = row != col
38
+ row = row[non_loop]
39
+ col = col[non_loop]
40
+
41
+ if len(row) == 0:
42
+ empty_ei = ops.zeros((2, 0), dtype=edge_index.dtype)
43
+ empty_ea = None if edge_attr is None else ops.zeros((0, *ops.shape(edge_attr)[1:]), dtype=edge_attr.dtype)
44
+ return empty_ei, empty_ea
45
+
46
+ edges = np.stack([row, col], axis=0)
47
+ # Coalesce duplicate edges
48
+ unique_edges, inv = np.unique(edges, axis=1, return_inverse=True)
49
+
50
+ out_edge_index = ops.convert_to_tensor(unique_edges, dtype=edge_index.dtype)
51
+
52
+ out_edge_attr = None
53
+ if edge_attr is not None:
54
+ ea_np = to_numpy(edge_attr)[non_loop]
55
+ num_unique = unique_edges.shape[1]
56
+ ea_tensor = ops.convert_to_tensor(ea_np, dtype=edge_attr.dtype)
57
+ inv_tensor = ops.convert_to_tensor(inv, dtype="int32")
58
+ if reduce == "sum":
59
+ out_edge_attr = segment_sum(ea_tensor, inv_tensor, num_segments=num_unique)
60
+ elif reduce == "mean":
61
+ sum_ea = segment_sum(ea_tensor, inv_tensor, num_segments=num_unique)
62
+ count = segment_sum(ops.ones_like(ea_tensor), inv_tensor, num_segments=num_unique)
63
+ out_edge_attr = sum_ea / ops.maximum(count, 1.0)
64
+ elif reduce == "max":
65
+ out_edge_attr = segment_max(ea_tensor, inv_tensor, num_segments=num_unique)
66
+
67
+ return out_edge_index, out_edge_attr
68
+
69
+
70
+ def pool_batch(perm, batch):
71
+ r"""Pools batch vector given representative indices `perm`.
72
+
73
+ Example:
74
+ ```python
75
+ import numpy as np
76
+ from k3_node.layers import pool_batch
77
+
78
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
79
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
80
+
81
+ perm = np.array([0, 3, 6, 9]) # nodes kept after pooling
82
+ print(tuple(pool_batch(perm, batch).shape)) # (4,): graph id of each kept node
83
+ ```
84
+ """
85
+ return ops.take(batch, perm, axis=0)
86
+
87
+
88
+ def pool_pos(cluster, pos):
89
+ r"""Pools node positions by computing average coordinates within each cluster.
90
+
91
+ Example:
92
+ ```python
93
+ import numpy as np
94
+ from k3_node.layers import pool_pos
95
+
96
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
97
+ pos = np.random.rand(10, 3).astype("float32") # 3D positions
98
+ cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
99
+
100
+ print(tuple(pool_pos(cluster, pos).shape)) # (5, 3): mean position of every cluster
101
+ ```
102
+ """
103
+ cluster = ops.cast(cluster, dtype="int32")
104
+ num_clusters = int(to_numpy(cluster).max()) + 1 if ops.shape(cluster)[0] > 0 else 0
105
+ sum_pos = segment_sum(pos, cluster, num_segments=num_clusters)
106
+ count = segment_sum(ops.ones_like(pos), cluster, num_segments=num_clusters)
107
+ return sum_pos / ops.maximum(count, 1.0)
108
+
109
+
110
+ def as_mutable_graph(data):
111
+ """Turns a (read-only) loader batch into a ``Data`` object that pooling can update in place.
112
+
113
+ ``ptr`` is dropped because it no longer matches the nodes once they are pooled.
114
+ """
115
+ if hasattr(data, "_asdict"):
116
+ from k3_node.data import Data
117
+
118
+ return Data(**{k: v for k, v in data._asdict().items() if k != "ptr" and v is not None})
119
+ return data
@@ -0,0 +1,174 @@
1
+ from typing import Callable, Optional, Tuple, Union
2
+ from keras import layers, ops
3
+
4
+ from .connect.filter_edges import FilterEdges
5
+ from .select.topk import SelectTopK
6
+ from k3_node.ops.segment import segment_max, segment_sum
7
+
8
+
9
+ class GraphConv(layers.Layer):
10
+ r"""Basic GraphConv layer for projection scoring in pooling layers."""
11
+ def __init__(
12
+ self,
13
+ in_channels: int,
14
+ out_channels: int,
15
+ aggr: str = "add",
16
+ bias: bool = True,
17
+ **kwargs,
18
+ ):
19
+ super().__init__(**kwargs)
20
+ self.in_channels = in_channels
21
+ self.out_channels = out_channels
22
+ self.aggr = aggr
23
+
24
+ self.lin_rel = layers.Dense(out_channels, use_bias=bias, name="lin_rel")
25
+ self.lin_root = layers.Dense(out_channels, use_bias=False, name="lin_root")
26
+
27
+ def reset_parameters(self):
28
+ pass
29
+
30
+ def build(self, input_shape=None):
31
+ if hasattr(self.lin_rel, "built") and not self.lin_rel.built:
32
+ self.lin_rel.build((None, self.in_channels))
33
+ if hasattr(self.lin_root, "built") and not self.lin_root.built:
34
+ self.lin_root.build((None, self.in_channels))
35
+ self.built = True
36
+
37
+ def call(self, x, edge_index, edge_weight: Optional[any] = None):
38
+ row = ops.cast(edge_index[0], dtype="int32")
39
+ col = ops.cast(edge_index[1], dtype="int32")
40
+ num_nodes = ops.shape(x)[0]
41
+
42
+ msg = ops.take(x, row, axis=0)
43
+ if edge_weight is not None:
44
+ msg = msg * ops.reshape(edge_weight, (-1, 1))
45
+
46
+ if self.aggr == "add":
47
+ aggr_out = segment_sum(msg, col, num_segments=num_nodes)
48
+ elif self.aggr == "mean":
49
+ sum_val = segment_sum(msg, col, num_segments=num_nodes)
50
+ count = segment_sum(ops.ones_like(msg), col, num_segments=num_nodes)
51
+ aggr_out = sum_val / ops.maximum(count, 1.0)
52
+ elif self.aggr == "max":
53
+ aggr_out = segment_max(msg, col, num_segments=num_nodes)
54
+ else:
55
+ aggr_out = segment_sum(msg, col, num_segments=num_nodes)
56
+
57
+ return self.lin_rel(aggr_out) + self.lin_root(x)
58
+
59
+
60
+ class SAGPooling(layers.Layer):
61
+ r"""The self-attention pooling operator from the `"Self-Attention Graph
62
+ Pooling" <https://arxiv.org/abs/1904.08082>`_ and `"Understanding
63
+ Attention and Generalization in Graph Neural Networks"
64
+ <https://arxiv.org/abs/1905.02850>`_ papers.
65
+
66
+ Example:
67
+ ```python
68
+ import numpy as np
69
+ from k3_node.layers import SAGPooling
70
+
71
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
72
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
73
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
74
+
75
+ layer = SAGPooling(in_channels=8, ratio=0.5) # keep half of the nodes of each graph
76
+ out = layer(x, edge_index, batch=batch)
77
+ x_pool, edge_index_pool, edge_attr_pool, batch_pool = out[0], out[1], out[2], out[3]
78
+ print(tuple(x_pool.shape)) # (6, 8)
79
+ ```
80
+ """
81
+ def __init__(
82
+ self,
83
+ in_channels: int,
84
+ ratio: Union[float, int] = 0.5,
85
+ GNN: Optional[any] = None,
86
+ min_score: Optional[float] = None,
87
+ multiplier: float = 1.0,
88
+ nonlinearity: Union[str, Callable] = "tanh",
89
+ **kwargs,
90
+ ):
91
+ super().__init__()
92
+
93
+ self.in_channels = in_channels
94
+ self.ratio = ratio
95
+ self.min_score = min_score
96
+ self.multiplier = multiplier
97
+ self.nonlinearity = nonlinearity
98
+
99
+ if GNN is None:
100
+ self.gnn = GraphConv(in_channels, 1, **kwargs)
101
+ elif callable(GNN):
102
+ self.gnn = GNN(in_channels, 1, **kwargs)
103
+ else:
104
+ self.gnn = GNN
105
+
106
+ self.select = SelectTopK(1, ratio, min_score, nonlinearity)
107
+ self.connect = FilterEdges()
108
+
109
+ def reset_parameters(self):
110
+ r"""Resets all learnable parameters of the module."""
111
+ if hasattr(self.gnn, "reset_parameters"):
112
+ self.gnn.reset_parameters()
113
+ self.select.reset_parameters()
114
+
115
+ def build(self, input_shape=None):
116
+ if hasattr(self.gnn, "built") and not self.gnn.built:
117
+ self.gnn.build(input_shape)
118
+ if hasattr(self.select, "built") and not self.select.built:
119
+ self.select.build(None)
120
+ if hasattr(self.connect, "built") and not self.connect.built:
121
+ self.connect.build(None)
122
+ self.built = True
123
+
124
+ def call(
125
+ self,
126
+ x,
127
+ edge_index,
128
+ edge_attr: Optional[any] = None,
129
+ batch: Optional[any] = None,
130
+ attn: Optional[any] = None,
131
+ ) -> Tuple[any, any, Optional[any], Optional[any], any, any]:
132
+ r"""Forward pass."""
133
+ num_nodes = ops.shape(x)[0]
134
+ if batch is None:
135
+ batch = ops.zeros((num_nodes,), dtype="int32")
136
+
137
+ if attn is None:
138
+ attn = x
139
+ if len(ops.shape(attn)) == 1:
140
+ attn = ops.expand_dims(attn, axis=-1)
141
+
142
+ attn = self.gnn(attn, edge_index)
143
+
144
+ select_out = self.select(attn, batch)
145
+
146
+ perm = select_out.node_index
147
+ score = select_out.weight
148
+
149
+ x_pooled = ops.take(x, perm, axis=0) * ops.expand_dims(score, axis=-1)
150
+ if self.multiplier != 1.0:
151
+ x_pooled = x_pooled * self.multiplier
152
+
153
+ connect_out = self.connect(select_out, edge_index, edge_attr, batch)
154
+
155
+ return (
156
+ x_pooled,
157
+ connect_out.edge_index,
158
+ connect_out.edge_attr,
159
+ connect_out.batch,
160
+ perm,
161
+ score,
162
+ )
163
+
164
+ def __repr__(self) -> str:
165
+ if self.min_score is None:
166
+ ratio = f"ratio={self.ratio}"
167
+ else:
168
+ ratio = f"min_score={self.min_score}"
169
+ gnn_name = self.gnn.__class__.__name__
170
+ return (
171
+ f"{self.__class__.__name__}({gnn_name}, {self.in_channels}, "
172
+ f"{ratio}, multiplier={self.multiplier})"
173
+ )
174
+
@@ -0,0 +1,10 @@
1
+ from .base import Select, SelectOutput
2
+ from .topk import SelectTopK, topk
3
+
4
+ __all__ = [
5
+ "SelectOutput",
6
+ "Select",
7
+ "topk",
8
+ "SelectTopK",
9
+ ]
10
+
@@ -0,0 +1,112 @@
1
+ from dataclasses import dataclass
2
+ from typing import Optional
3
+ from keras import layers, ops
4
+
5
+
6
+ @dataclass
7
+ class SelectOutput:
8
+ r"""The output of the :class:`Select` method, which holds an assignment
9
+ from selected nodes to their respective cluster(s).
10
+
11
+ Args:
12
+ node_index: The indices of the selected nodes.
13
+ num_nodes: The number of nodes.
14
+ cluster_index: The indices of the clusters each node in
15
+ :obj:`node_index` is assigned to.
16
+ num_clusters: The number of clusters.
17
+ weight (optional): A weight vector, denoting the strength
18
+ of the assignment of a node to its cluster. (default: :obj:`None`)
19
+ """
20
+ node_index: any
21
+ num_nodes: int
22
+ cluster_index: any
23
+ num_clusters: int
24
+ weight: Optional[any] = None
25
+
26
+ def __post_init__(self):
27
+ shape_node = getattr(self.node_index, "shape", None)
28
+ shape_cluster = getattr(self.cluster_index, "shape", None)
29
+ if shape_node is not None and len(shape_node) != 1:
30
+ raise ValueError(
31
+ f"Expected 'node_index' to be one-dimensional "
32
+ f"(got {len(shape_node)} dimensions)"
33
+ )
34
+ if shape_cluster is not None and len(shape_cluster) != 1:
35
+ raise ValueError(
36
+ f"Expected 'cluster_index' to be one-dimensional "
37
+ f"(got {len(shape_cluster)} dimensions)"
38
+ )
39
+ if (
40
+ shape_node is not None
41
+ and shape_cluster is not None
42
+ and len(shape_node) > 0
43
+ and len(shape_cluster) > 0
44
+ and shape_node[0] is not None
45
+ and shape_cluster[0] is not None
46
+ and shape_node[0] != shape_cluster[0]
47
+ ):
48
+ raise ValueError(
49
+ f"Expected 'node_index' and 'cluster_index' to hold the same "
50
+ f"number of values (got {shape_node[0]} and "
51
+ f"{shape_cluster[0]} values)"
52
+ )
53
+ if self.weight is not None:
54
+ shape_weight = getattr(self.weight, "shape", None)
55
+ if shape_weight is not None and len(shape_weight) != 1:
56
+ raise ValueError(
57
+ f"Expected 'weight' vector to be one-dimensional "
58
+ f"(got {len(shape_weight)} dimensions)"
59
+ )
60
+ if (
61
+ shape_weight is not None
62
+ and shape_node is not None
63
+ and len(shape_weight) > 0
64
+ and len(shape_node) > 0
65
+ and shape_weight[0] is not None
66
+ and shape_node[0] is not None
67
+ and shape_weight[0] != shape_node[0]
68
+ ):
69
+ raise ValueError(
70
+ f"Expected 'weight' to hold {shape_node[0]} "
71
+ f"values (got {shape_weight[0]} values)"
72
+ )
73
+
74
+
75
+ import keras
76
+
77
+ if keras.config.backend() == "jax":
78
+ try:
79
+ import jax
80
+ from jax.tree_util import register_pytree_node
81
+
82
+ register_pytree_node(
83
+ SelectOutput,
84
+ lambda s: (
85
+ (s.node_index, s.cluster_index, s.weight),
86
+ (s.num_nodes, s.num_clusters),
87
+ ),
88
+ lambda aux, children: SelectOutput(
89
+ children[0], aux[0], children[1], aux[1], children[2]
90
+ ),
91
+ )
92
+ except Exception:
93
+ pass
94
+
95
+
96
+
97
+ class Select(layers.Layer):
98
+ r"""An abstract base class for implementing custom node selections as
99
+ described in the `"Understanding Pooling in Graph Neural Networks"
100
+ <https://arxiv.org/abs/1905.05178>`_ paper, which maps the nodes of an
101
+ input graph to supernodes in the coarsened graph.
102
+ """
103
+ def reset_parameters(self):
104
+ r"""Resets all learnable parameters of the module."""
105
+ pass
106
+
107
+ def call(self, *args, **kwargs) -> SelectOutput:
108
+ raise NotImplementedError
109
+
110
+ def __repr__(self) -> str:
111
+ return f'{self.__class__.__name__}()'
112
+
@@ -0,0 +1,206 @@
1
+ import numpy as np
2
+ from typing import Callable, Optional, Union
3
+ from keras import initializers, layers, ops
4
+
5
+ from .base import Select, SelectOutput
6
+
7
+
8
+ from k3_node.layers.conv.utils import is_tracing
9
+ from k3_node.ops.segment import segment_max, segment_sum
10
+ from k3_node.ops.creation import full
11
+
12
+
13
+ def topk(
14
+ x,
15
+ ratio: Optional[Union[float, int]],
16
+ batch,
17
+ min_score: Optional[float] = None,
18
+ tol: float = 1e-7,
19
+ ):
20
+ r"""Selects top-k items according to score and batch assignment.
21
+
22
+ Example:
23
+ ```python
24
+ import numpy as np
25
+ from k3_node.layers import topk
26
+
27
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
28
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
29
+
30
+ score = np.random.rand(10).astype("float32")
31
+ perm = topk(score, ratio=0.5, batch=batch) # indices of the top 50% nodes per graph
32
+ print(tuple(perm.shape)) # (6,)
33
+ ```
34
+ """
35
+ # Plain NumPy inputs cannot be mixed with backend tensors; convert them first.
36
+ x = ops.convert_to_tensor(x) if isinstance(x, np.ndarray) else x
37
+ batch = ops.convert_to_tensor(batch) if isinstance(batch, np.ndarray) else batch
38
+ if is_tracing(x) or is_tracing(batch):
39
+ return ops.arange(ops.shape(x)[0], dtype="int32")
40
+
41
+ batch = ops.cast(batch, dtype="int32")
42
+ num_nodes = ops.shape(x)[0]
43
+
44
+ num_graphs = ops.max(batch) + 1 if ops.shape(batch)[0] > 0 else 0
45
+ try:
46
+ num_graphs = int(num_graphs)
47
+ except (TypeError, ValueError):
48
+ pass
49
+
50
+ if min_score is not None:
51
+ scores_max = segment_max(x, batch, num_segments=num_graphs)
52
+ scores_max_expanded = ops.take(scores_max, batch, axis=0) - tol
53
+ scores_min = ops.minimum(scores_max_expanded, min_score)
54
+ mask = x > scores_min
55
+ perm = ops.where(mask)
56
+ if isinstance(perm, (tuple, list)):
57
+ perm = perm[0]
58
+ perm = ops.reshape(perm, (-1,))
59
+ return ops.cast(perm, "int32")
60
+
61
+ if ratio is not None:
62
+ ones = ops.ones((num_nodes,), dtype="int32")
63
+ num_nodes_per_graph = segment_sum(ones, batch, num_segments=num_graphs)
64
+
65
+ if ratio >= 1:
66
+ k = full(ops.shape(num_nodes_per_graph), int(ratio), dtype="int32")
67
+ else:
68
+ k = ops.cast(
69
+ ops.ceil(ratio * ops.cast(num_nodes_per_graph, x.dtype)),
70
+ dtype="int32",
71
+ )
72
+
73
+ # Composite key: sorts by batch ascending, then by score descending
74
+ score_span = ops.max(x) - ops.min(x) + 1.0
75
+ key = ops.cast(batch, x.dtype) * (score_span * 2.0) - x
76
+ perm = ops.argsort(key)
77
+
78
+ batch_sorted = ops.take(batch, perm, axis=0)
79
+ # ptr for cumsum
80
+ ptr = ops.concatenate(
81
+ [ops.zeros((1,), dtype="int32"), ops.cumsum(num_nodes_per_graph)[:-1]],
82
+ axis=0,
83
+ )
84
+ rank_in_graph = ops.arange(num_nodes, dtype="int32") - ops.take(ptr, batch_sorted, axis=0)
85
+ mask = rank_in_graph < ops.take(k, batch_sorted, axis=0)
86
+ valid_idx = ops.where(mask)
87
+ if isinstance(valid_idx, (tuple, list)):
88
+ valid_idx = valid_idx[0]
89
+ valid_idx = ops.reshape(valid_idx, (-1,))
90
+ return ops.cast(ops.take(perm, valid_idx, axis=0), "int32")
91
+
92
+ raise ValueError("At least one of 'ratio' and 'min_score' must be specified.")
93
+
94
+
95
+ class SelectTopK(Select):
96
+ r"""Selects the top-:math:`k` nodes with highest projection scores.
97
+
98
+ Example:
99
+ ```python
100
+ import numpy as np
101
+ from k3_node.layers import SelectTopK
102
+
103
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
104
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
105
+
106
+ select = SelectTopK(in_channels=8, ratio=0.5)
107
+ out = select(x, batch)
108
+ print(tuple(out.node_index.shape)) # (6,): indices of the kept nodes
109
+ print(tuple(out.weight.shape)) # (6,): their scores
110
+ ```
111
+ """
112
+ def __init__(
113
+ self,
114
+ in_channels: int,
115
+ ratio: Union[int, float] = 0.5,
116
+ min_score: Optional[float] = None,
117
+ act: Union[str, Callable] = "tanh",
118
+ **kwargs,
119
+ ):
120
+ super().__init__(**kwargs)
121
+
122
+ if ratio is None and min_score is None:
123
+ raise ValueError(
124
+ f"At least one of 'ratio' and 'min_score' must be specified in '{self.__class__.__name__}'"
125
+ )
126
+
127
+ self.in_channels = in_channels
128
+ self.ratio = ratio
129
+ self.min_score = min_score
130
+ self.act_fn = act if callable(act) else layers.Activation(act)
131
+
132
+ self.weight = self.add_weight(
133
+ shape=(1, in_channels),
134
+ initializer=initializers.RandomUniform(
135
+ minval=-1.0 / (in_channels**0.5), maxval=1.0 / (in_channels**0.5)
136
+ ),
137
+ trainable=True,
138
+ name="weight",
139
+ )
140
+
141
+ def reset_parameters(self):
142
+ limit = 1.0 / (self.in_channels**0.5)
143
+ init = initializers.RandomUniform(minval=-limit, maxval=limit)
144
+ self.weight.assign(init(self.weight.shape, dtype=self.weight.dtype))
145
+
146
+ def build(self, input_shape=None):
147
+ if hasattr(self.act_fn, "built") and not self.act_fn.built:
148
+ self.act_fn.build(input_shape)
149
+ self.built = True
150
+
151
+ def call(self, x, batch=None) -> SelectOutput:
152
+ num_nodes = ops.shape(x)[0]
153
+ if batch is None:
154
+ batch = ops.zeros((num_nodes,), dtype="int32")
155
+ else:
156
+ batch = ops.cast(batch, dtype="int32")
157
+
158
+ if len(ops.shape(x)) == 1:
159
+ x = ops.expand_dims(x, axis=-1)
160
+
161
+ score = ops.sum(x * self.weight, axis=-1)
162
+
163
+ if is_tracing(x) or is_tracing(batch):
164
+ node_index = ops.arange(num_nodes, dtype="int32")
165
+ return SelectOutput(
166
+ node_index=node_index,
167
+ num_nodes=num_nodes,
168
+ cluster_index=node_index,
169
+ num_clusters=num_nodes,
170
+ weight=score,
171
+ )
172
+
173
+ if self.min_score is None:
174
+ norm_w = ops.sqrt(ops.sum(ops.power(self.weight, 2), axis=-1))
175
+ score = self.act_fn(score / norm_w)
176
+ else:
177
+ # Graph-wise softmax
178
+ num_graphs = ops.max(batch) + 1 if ops.shape(batch)[0] > 0 else 0
179
+ try:
180
+ num_graphs = int(num_graphs)
181
+ except (TypeError, ValueError):
182
+ pass
183
+ score_max = segment_max(score, batch, num_segments=num_graphs)
184
+ score_max_exp = ops.take(score_max, batch, axis=0)
185
+ exp_score = ops.exp(score - score_max_exp)
186
+ exp_sum = segment_sum(exp_score, batch, num_segments=num_graphs)
187
+ exp_sum_exp = ops.take(exp_sum, batch, axis=0)
188
+ score = exp_score / (exp_sum_exp + 1e-12)
189
+
190
+ node_index = topk(score, self.ratio, batch, self.min_score)
191
+ num_selected = ops.shape(node_index)[0]
192
+
193
+ return SelectOutput(
194
+ node_index=node_index,
195
+ num_nodes=num_nodes,
196
+ cluster_index=ops.arange(num_selected, dtype="int32"),
197
+ num_clusters=num_selected,
198
+ weight=ops.take(score, node_index, axis=0),
199
+ )
200
+
201
+ def __repr__(self) -> str:
202
+ if self.min_score is None:
203
+ arg = f"ratio={self.ratio}"
204
+ else:
205
+ arg = f"min_score={self.min_score}"
206
+ return f"{self.__class__.__name__}({self.in_channels}, {arg})"