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,165 @@
1
+ from typing import Callable, Optional, Tuple
2
+ from keras import ops
3
+
4
+ from .consecutive import consecutive_cluster
5
+ from .pool import as_mutable_graph, pool_batch, pool_edge, pool_pos
6
+ from k3_node.ops.segment import segment_sum
7
+ from k3_node.ops.host import to_numpy
8
+
9
+
10
+ def _avg_pool_x(cluster, x, size: Optional[int] = None):
11
+ cluster = ops.cast(cluster, dtype="int32")
12
+ if size is None:
13
+ size = int(to_numpy(cluster).max()) + 1 if ops.shape(cluster)[0] > 0 else 0
14
+ sum_x = segment_sum(x, cluster, num_segments=size)
15
+ ones = ops.ones_like(x)
16
+ count = segment_sum(ones, cluster, num_segments=size)
17
+ return sum_x / ops.maximum(count, 1.0)
18
+
19
+
20
+ def avg_pool_x(
21
+ cluster,
22
+ x,
23
+ batch,
24
+ batch_size: Optional[int] = None,
25
+ size: Optional[int] = None,
26
+ ) -> Tuple[any, Optional[any]]:
27
+ r"""Average-pools node features according to the clustering defined in `cluster`.
28
+
29
+ Example:
30
+ ```python
31
+ import numpy as np
32
+ from k3_node.layers import avg_pool_x
33
+
34
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
35
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
36
+ cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
37
+
38
+ x_pool, batch_pool = avg_pool_x(cluster, x, batch)
39
+ print(tuple(x_pool.shape)) # (5, 8)
40
+ ```
41
+ """
42
+ if size is not None:
43
+ if batch_size is None:
44
+ batch_size = int(to_numpy(batch).max()) + 1
45
+ return _avg_pool_x(cluster, x, batch_size * size), None
46
+
47
+ cluster, perm = consecutive_cluster(cluster)
48
+ x = _avg_pool_x(cluster, x)
49
+ batch = pool_batch(perm, batch)
50
+ return x, batch
51
+
52
+
53
+ def avg_pool(
54
+ cluster,
55
+ data,
56
+ transform: Optional[Callable] = None,
57
+ edge_index: Optional[any] = None,
58
+ edge_attr: Optional[any] = None,
59
+ batch: Optional[any] = None,
60
+ pos: Optional[any] = None,
61
+ ):
62
+ r"""Pools and coarsens a graph given by `data` according to `cluster` using averaging.
63
+
64
+ Example:
65
+ ```python
66
+ import numpy as np
67
+ from k3_node.layers import avg_pool
68
+
69
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
70
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
71
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
72
+ cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
73
+
74
+ x_pool, edge_index_pool, batch_pool = avg_pool(cluster, x, edge_index, batch=batch)
75
+ print(tuple(x_pool.shape)) # (5, 8)
76
+ ```
77
+ """
78
+ cluster, perm = consecutive_cluster(cluster)
79
+
80
+ if hasattr(data, "x"):
81
+ data = as_mutable_graph(data)
82
+ x = getattr(data, "x", None)
83
+ if x is not None:
84
+ data.x = _avg_pool_x(cluster, x)
85
+
86
+ edge_index = getattr(data, "edge_index", None)
87
+ edge_attr = getattr(data, "edge_attr", None)
88
+ if edge_index is not None:
89
+ data.edge_index, data.edge_attr = pool_edge(cluster, edge_index, edge_attr, reduce="mean")
90
+
91
+ batch = getattr(data, "batch", None)
92
+ if batch is not None:
93
+ data.batch = pool_batch(perm, batch)
94
+
95
+ pos = getattr(data, "pos", None)
96
+ if pos is not None:
97
+ data.pos = pool_pos(cluster, pos)
98
+
99
+ if transform is not None:
100
+ data = transform(data)
101
+
102
+ return data
103
+
104
+ # Raw tensor mode
105
+ pooled_x = _avg_pool_x(cluster, data)
106
+ pooled_edge_index, pooled_edge_attr = (None, None)
107
+ if edge_index is not None:
108
+ pooled_edge_index, pooled_edge_attr = pool_edge(cluster, edge_index, edge_attr, reduce="mean")
109
+ pooled_batch = pool_batch(perm, batch) if batch is not None else None
110
+
111
+ return pooled_x, pooled_edge_index, pooled_batch
112
+
113
+
114
+ def avg_pool_neighbor_x(
115
+ data,
116
+ edge_index=None,
117
+ flow: str = "source_to_target",
118
+ ):
119
+ r"""Average-pools neighboring node features.
120
+
121
+ Example:
122
+ ```python
123
+ import numpy as np
124
+ from k3_node.layers import avg_pool_neighbor_x
125
+
126
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
127
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
128
+
129
+ out = avg_pool_neighbor_x(x, edge_index=edge_index) # pool each node with its neighbors
130
+ print(tuple(out.shape)) # (10, 8)
131
+ ```
132
+ """
133
+ if hasattr(data, "x"):
134
+ x = data.x
135
+ edge_index = data.edge_index
136
+ is_data_obj = True
137
+ else:
138
+ x = data
139
+ is_data_obj = False
140
+ if edge_index is None:
141
+ raise ValueError("edge_index must be provided if data is a tensor.")
142
+
143
+ num_nodes = getattr(data, "num_nodes", None) if is_data_obj else ops.shape(x)[0]
144
+ if num_nodes is None:
145
+ num_nodes = ops.shape(x)[0]
146
+
147
+ # Add self-loops
148
+ loop_idx = ops.arange(num_nodes, dtype=edge_index.dtype)
149
+ loop_edge = ops.stack([loop_idx, loop_idx], axis=0)
150
+ full_edge_index = ops.concatenate([edge_index, loop_edge], axis=1)
151
+
152
+ row = full_edge_index[0]
153
+ col = full_edge_index[1]
154
+ row, col = (row, col) if flow == "source_to_target" else (col, row)
155
+
156
+ col = ops.cast(col, dtype="int32")
157
+ x_src = ops.take(x, row, axis=0)
158
+ sum_x = segment_sum(x_src, col, num_segments=num_nodes)
159
+ ones = ops.ones_like(x_src)
160
+ count = segment_sum(ones, col, num_segments=num_nodes)
161
+ out_x = sum_x / ops.maximum(count, 1.0)
162
+ if is_data_obj:
163
+ data.x = out_x
164
+ return data
165
+ return out_x
@@ -0,0 +1,168 @@
1
+ from typing import NamedTuple, Optional, Tuple
2
+ from keras import layers, ops
3
+ import numpy as np
4
+ from k3_node.ops.segment import segment_sum
5
+
6
+
7
+ class UnpoolInfo(NamedTuple):
8
+ edge_index: any
9
+ cluster: any
10
+ batch: any
11
+
12
+
13
+ class ClusterPooling(layers.Layer):
14
+ r"""The cluster pooling operator from the `"Edge-Based Graph Component
15
+ Pooling" <https://arxiv.org/abs/2409.11856>`_ paper.
16
+
17
+ Example:
18
+ ```python
19
+ import numpy as np
20
+ from k3_node.layers import ClusterPooling
21
+
22
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
23
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
24
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
25
+
26
+ layer = ClusterPooling(in_channels=8)
27
+ x_pool, edge_index_pool, batch_pool, unpool_info = layer(x, edge_index, batch)
28
+ # The number of clusters depends on the learned edge scores
29
+ print(x_pool.shape[0] <= 10, x_pool.shape[1]) # True 8: fewer nodes, same features
30
+ ```
31
+ """
32
+ def __init__(
33
+ self,
34
+ in_channels: int,
35
+ edge_score_method: str = "tanh",
36
+ dropout: float = 0.0,
37
+ threshold: Optional[float] = None,
38
+ **kwargs,
39
+ ):
40
+ super().__init__(**kwargs)
41
+ assert edge_score_method in ["tanh", "sigmoid", "log_softmax"]
42
+
43
+ if threshold is None:
44
+ threshold = 0.5 if edge_score_method == "sigmoid" else 0.0
45
+
46
+ self.in_channels = in_channels
47
+ self.edge_score_method = edge_score_method
48
+ self.dropout_rate = dropout
49
+ self.drop = layers.Dropout(dropout) if dropout > 0.0 else None
50
+ self.threshold = threshold
51
+
52
+ self.lin = layers.Dense(1, use_bias=True, name="lin")
53
+
54
+ def build(self, input_shape=None):
55
+ self.lin.build((None, 2 * self.in_channels))
56
+ super().build(input_shape)
57
+
58
+ def reset_parameters(self):
59
+ self.lin.reset_parameters()
60
+
61
+ def call(
62
+ self,
63
+ x,
64
+ edge_index,
65
+ batch,
66
+ training: bool = False,
67
+ ) -> Tuple[any, any, any, UnpoolInfo]:
68
+ r"""Forward pass."""
69
+ from k3_node.layers.conv.utils import eager_only_placeholder
70
+ if eager_only_placeholder("ClusterPooling", x, edge_index):
71
+ num_nodes = ops.shape(x)[0]
72
+ unpool_info = UnpoolInfo(edge_index, ops.arange(num_nodes, dtype="int32"), batch)
73
+ return x, edge_index, batch, unpool_info
74
+
75
+ edge_index_np = ops.convert_to_numpy(edge_index).astype(np.int64)
76
+ mask = edge_index_np[0] != edge_index_np[1]
77
+ edge_index_filtered = edge_index_np[:, mask]
78
+
79
+ row = ops.convert_to_tensor(edge_index_filtered[0], dtype="int32")
80
+ col = ops.convert_to_tensor(edge_index_filtered[1], dtype="int32")
81
+
82
+ edge_attr = ops.concatenate([ops.take(x, row, axis=0), ops.take(x, col, axis=0)], axis=-1)
83
+ edge_score = ops.reshape(self.lin(edge_attr), (-1,))
84
+ if self.drop is not None:
85
+ edge_score = self.drop(edge_score, training=training)
86
+
87
+ if self.edge_score_method == "tanh":
88
+ edge_score = ops.tanh(edge_score)
89
+ elif self.edge_score_method == "sigmoid":
90
+ edge_score = ops.sigmoid(edge_score)
91
+ else:
92
+ edge_score = ops.log_softmax(edge_score, axis=0)
93
+
94
+ edge_index_tensor = ops.convert_to_tensor(edge_index_filtered, dtype=edge_index.dtype)
95
+ return self._merge_edges(x, edge_index_tensor, batch, edge_score)
96
+
97
+ def _merge_edges(
98
+ self,
99
+ x,
100
+ edge_index,
101
+ batch,
102
+ edge_score,
103
+ ) -> Tuple[any, any, any, UnpoolInfo]:
104
+ from scipy.sparse import coo_matrix
105
+ from scipy.sparse.csgraph import connected_components
106
+
107
+ from k3_node.layers.conv.utils import host_callback
108
+
109
+ num_nodes = int(ops.shape(x)[0])
110
+ num_edges = int(ops.shape(edge_index)[1])
111
+ threshold = self.threshold
112
+
113
+ def contract(edge_index_np, edge_score_np, batch_np):
114
+ # Clusters are the weakly connected components of the edges scoring above the threshold.
115
+ edge_index_np = edge_index_np.astype(np.int64)
116
+ edge_contract = edge_index_np[:, edge_score_np > threshold]
117
+ if edge_contract.shape[1] > 0:
118
+ adj = coo_matrix(
119
+ (np.ones(edge_contract.shape[1]), (edge_contract[0], edge_contract[1])),
120
+ shape=(num_nodes, num_nodes),
121
+ )
122
+ _, cluster_np = connected_components(adj, directed=True, connection="weak")
123
+ else:
124
+ cluster_np = np.arange(num_nodes)
125
+ num_clusters = int(np.max(cluster_np)) + 1 if num_nodes > 0 else 0
126
+
127
+ # Nodes without any contracted edge keep their own features (unit diagonal score).
128
+ single = np.ones(num_nodes, dtype=bool)
129
+ single[edge_contract[0]] = False
130
+ single[edge_contract[1]] = False
131
+
132
+ # Coarsened edges between distinct clusters, in (row, col) order.
133
+ pairs = cluster_np[edge_index_np]
134
+ pairs = np.unique(pairs[:, pairs[0] != pairs[1]], axis=1)
135
+ edges_pad = np.zeros((2, num_edges), dtype=np.int64)
136
+ edges_pad[:, : pairs.shape[1]] = pairs
137
+
138
+ batch_pad = np.zeros(num_nodes, dtype=np.int64)
139
+ batch_pad[cluster_np] = batch_np
140
+ return cluster_np, num_clusters, single, edges_pad, pairs.shape[1], batch_pad
141
+
142
+ cluster, num_clusters, single, edges_pad, num_new_edges, batch_pad = host_callback(
143
+ contract,
144
+ [((num_nodes,), "int32"), ((), "int32"), ((num_nodes,), "float32"),
145
+ ((2, num_edges), "int32"), ((), "int32"), ((num_nodes,), "int32")],
146
+ edge_index, edge_score, batch,
147
+ )
148
+ num_clusters, num_new_edges = int(num_clusters), int(num_new_edges)
149
+
150
+ # x_out = (S @ C)^T @ x, computed sparsely: every edge (row -> col) adds score * x[row] to
151
+ # cluster(col), and every unmatched node adds its own features to its cluster. The score
152
+ # enters as a tensor so gradients reach the scoring layer.
153
+ row = ops.cast(edge_index[0], "int32")
154
+ col = ops.cast(edge_index[1], "int32")
155
+ msgs = ops.expand_dims(ops.cast(edge_score, x.dtype), -1) * ops.take(x, row, axis=0)
156
+ x_out = segment_sum(msgs, ops.take(cluster, col, axis=0), num_segments=num_clusters)
157
+ x_out = x_out + segment_sum(
158
+ x * ops.expand_dims(ops.cast(single, x.dtype), -1), cluster, num_segments=num_clusters
159
+ )
160
+
161
+ edge_index_out = ops.cast(edges_pad[:, :num_new_edges], edge_index.dtype)
162
+ batch_out = ops.cast(batch_pad[:num_clusters], batch.dtype)
163
+
164
+ unpool_info = UnpoolInfo(edge_index, cluster, batch)
165
+ return x_out, edge_index_out, batch_out, unpool_info
166
+
167
+ def __repr__(self) -> str:
168
+ return f"{self.__class__.__name__}({self.in_channels})"
@@ -0,0 +1,10 @@
1
+ from .base import Connect, ConnectOutput
2
+ from .filter_edges import FilterEdges, filter_adj
3
+
4
+ __all__ = [
5
+ "ConnectOutput",
6
+ "Connect",
7
+ "filter_adj",
8
+ "FilterEdges",
9
+ ]
10
+
@@ -0,0 +1,103 @@
1
+ from dataclasses import dataclass
2
+ from typing import Optional
3
+ from keras import layers, ops
4
+
5
+ from ..select.base import SelectOutput
6
+
7
+
8
+ @dataclass
9
+ class ConnectOutput:
10
+ r"""The output of the :class:`Connect` method, which holds the coarsened
11
+ graph structure, and optional pooled edge features and batch vectors.
12
+
13
+ Args:
14
+ edge_index: The edge indices of the coarsened graph.
15
+ edge_attr: The pooled edge features of the coarsened graph. (default: None)
16
+ batch: The pooled batch vector of the coarsened graph. (default: None)
17
+ """
18
+ edge_index: any
19
+ edge_attr: Optional[any] = None
20
+ batch: Optional[any] = None
21
+
22
+ def __post_init__(self):
23
+ shape_edge = getattr(self.edge_index, "shape", None)
24
+ if shape_edge is not None:
25
+ if len(shape_edge) != 2:
26
+ raise ValueError(
27
+ f"Expected 'edge_index' to be two-dimensional "
28
+ f"(got {len(shape_edge)} dimensions)"
29
+ )
30
+ if shape_edge[0] is not None and shape_edge[0] != 2:
31
+ raise ValueError(
32
+ f"Expected 'edge_index' to have size '2' in the first dimension "
33
+ f"(got '{shape_edge[0]}')"
34
+ )
35
+ if self.edge_attr is not None:
36
+ shape_attr = getattr(self.edge_attr, "shape", None)
37
+ if (
38
+ shape_edge is not None
39
+ and shape_attr is not None
40
+ and len(shape_edge) == 2
41
+ and len(shape_attr) >= 1
42
+ and shape_edge[1] is not None
43
+ and shape_attr[0] is not None
44
+ and shape_attr[0] != shape_edge[1]
45
+ ):
46
+ raise ValueError(
47
+ f"Expected 'edge_index' and 'edge_attr' to hold the same number "
48
+ f"of edges (got {shape_edge[1]} and {shape_attr[0]} edges)"
49
+ )
50
+
51
+
52
+ import keras
53
+
54
+ if keras.config.backend() == "jax":
55
+ try:
56
+ import jax
57
+ from jax.tree_util import register_pytree_node
58
+
59
+ register_pytree_node(
60
+ ConnectOutput,
61
+ lambda c: ((c.edge_index, c.edge_attr, c.batch), ()),
62
+ lambda aux, children: ConnectOutput(children[0], children[1], children[2]),
63
+ )
64
+ except Exception:
65
+ pass
66
+
67
+
68
+
69
+ class Connect(layers.Layer):
70
+ r"""An abstract base class for implementing custom edge connection
71
+ operators as described in the `"Understanding Pooling in Graph Neural
72
+ Networks" <https://arxiv.org/abs/1905.05178>`_ paper.
73
+ """
74
+ def reset_parameters(self):
75
+ r"""Resets all learnable parameters of the module."""
76
+ pass
77
+
78
+ def __call__(self, *args, **kwargs):
79
+ if len(args) > 0 and isinstance(args[0], SelectOutput):
80
+ return self.call(*args, **kwargs)
81
+ return super().__call__(*args, **kwargs)
82
+
83
+ def call(
84
+ self,
85
+ select_output: SelectOutput,
86
+ edge_index,
87
+ edge_attr: Optional[any] = None,
88
+ batch: Optional[any] = None,
89
+ ) -> ConnectOutput:
90
+ raise NotImplementedError
91
+
92
+ @staticmethod
93
+ def get_pooled_batch(
94
+ select_output: SelectOutput,
95
+ batch: Optional[any],
96
+ ) -> Optional[any]:
97
+ r"""Returns the batch vector of the coarsened graph."""
98
+ if batch is None:
99
+ return None
100
+ return ops.take(batch, select_output.node_index, axis=0)
101
+
102
+ def __repr__(self) -> str:
103
+ return f'{self.__class__.__name__}()'
@@ -0,0 +1,113 @@
1
+ from typing import Optional, Tuple
2
+ from keras import ops
3
+ import numpy as np
4
+
5
+ from .base import Connect, ConnectOutput
6
+ from ..select.base import SelectOutput
7
+
8
+
9
+ from k3_node.layers.conv.utils import is_tracing
10
+ from k3_node.ops.creation import full
11
+
12
+
13
+ def filter_adj(
14
+ edge_index,
15
+ edge_attr: Optional[any] = None,
16
+ node_index=None,
17
+ cluster_index: Optional[any] = None,
18
+ num_nodes: Optional[int] = None,
19
+ ) -> Tuple[any, Optional[any]]:
20
+ r"""Filters out edges if their incident nodes are not in any cluster.
21
+
22
+ Example:
23
+ ```python
24
+ import numpy as np
25
+ from k3_node.layers import filter_adj
26
+
27
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
28
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
29
+
30
+ kept_nodes = np.array([0, 2, 4, 6, 8])
31
+ edge_index_new, _ = filter_adj(edge_index, node_index=kept_nodes) # edges between kept nodes, relabeled
32
+ print(edge_index_new.shape[0]) # 2
33
+ ```
34
+ """
35
+ if is_tracing(edge_index) or (node_index is not None and is_tracing(node_index)):
36
+ return edge_index, edge_attr
37
+
38
+ if node_index is None:
39
+ return edge_index, edge_attr
40
+
41
+ edge_index = ops.cast(edge_index, "int32")
42
+ node_index = ops.cast(node_index, "int32")
43
+ if cluster_index is None:
44
+ cluster_index = ops.arange(ops.shape(node_index)[0], dtype="int32")
45
+ else:
46
+ cluster_index = ops.cast(cluster_index, "int32")
47
+
48
+ if num_nodes is None:
49
+ num_nodes = ops.max(node_index) + 1 if ops.shape(node_index)[0] > 0 else 0
50
+ if ops.shape(edge_index)[1] > 0:
51
+ num_nodes = ops.maximum(num_nodes, ops.max(edge_index) + 1)
52
+ try:
53
+ num_nodes = int(num_nodes)
54
+ except (TypeError, ValueError):
55
+ pass
56
+
57
+ mapping = full((num_nodes,), -1, dtype="int32")
58
+ mapping = ops.scatter_update(mapping, ops.expand_dims(node_index, -1), cluster_index)
59
+
60
+ row = ops.take(mapping, edge_index[0], axis=0)
61
+ col = ops.take(mapping, edge_index[1], axis=0)
62
+ mask = (row >= 0) & (col >= 0)
63
+ valid_idx = ops.where(mask)
64
+ if isinstance(valid_idx, (tuple, list)):
65
+ valid_idx = valid_idx[0]
66
+ valid_idx = ops.reshape(valid_idx, (-1,))
67
+
68
+ new_edge_index = ops.stack(
69
+ [ops.take(row, valid_idx, axis=0), ops.take(col, valid_idx, axis=0)], axis=0
70
+ )
71
+
72
+ new_edge_attr = None
73
+ if edge_attr is not None:
74
+ new_edge_attr = ops.take(edge_attr, valid_idx, axis=0)
75
+
76
+ return new_edge_index, new_edge_attr
77
+
78
+
79
+ class FilterEdges(Connect):
80
+ r"""Filters out edges if their incident nodes are not in any cluster.
81
+
82
+ Example:
83
+ ```python
84
+ import numpy as np
85
+ from k3_node.layers import FilterEdges, SelectTopK
86
+
87
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
88
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
89
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
90
+
91
+ select_output = SelectTopK(in_channels=8, ratio=0.5)(x, batch)
92
+ connect = FilterEdges()
93
+ out = connect(select_output, edge_index, batch=batch) # keep edges between kept nodes
94
+ print(out.edge_index.shape[0]) # 2
95
+ ```
96
+ """
97
+ def call(
98
+ self,
99
+ select_output: SelectOutput,
100
+ edge_index,
101
+ edge_attr: Optional[any] = None,
102
+ batch: Optional[any] = None,
103
+ ) -> ConnectOutput:
104
+ new_edge_index, new_edge_attr = filter_adj(
105
+ edge_index,
106
+ edge_attr,
107
+ select_output.node_index,
108
+ select_output.cluster_index,
109
+ num_nodes=select_output.num_nodes,
110
+ )
111
+ new_batch = self.get_pooled_batch(select_output, batch)
112
+ return ConnectOutput(new_edge_index, new_edge_attr, new_batch)
113
+
@@ -0,0 +1,30 @@
1
+ from typing import Tuple
2
+ from keras import ops
3
+ import numpy as np
4
+ from k3_node.ops.host import to_numpy
5
+
6
+
7
+ def consecutive_cluster(src) -> Tuple[any, any]:
8
+ r"""Maps elements in `src` to consecutive integers starting from 0,
9
+ returning the mapped indices and a permutation array of representative indices.
10
+
11
+ Example:
12
+ ```python
13
+ import numpy as np
14
+ from k3_node.layers import consecutive_cluster
15
+
16
+ cluster = np.array([4, 4, 9, 2, 9])
17
+ new_cluster, perm = consecutive_cluster(cluster) # relabel cluster ids to 0..num_clusters-1
18
+ print(tuple(new_cluster.shape), tuple(perm.shape)) # (5,) (3,)
19
+ ```
20
+ """
21
+ src_np = to_numpy(src)
22
+ unique, inv = np.unique(src_np, return_inverse=True)
23
+ perm = np.empty(len(unique), dtype=inv.dtype)
24
+ arange = np.arange(len(inv), dtype=inv.dtype)
25
+ perm[inv] = arange
26
+
27
+ inv_tensor = ops.convert_to_tensor(inv, dtype=src.dtype)
28
+ perm_tensor = ops.convert_to_tensor(perm, dtype=src.dtype)
29
+ return inv_tensor, perm_tensor
30
+
@@ -0,0 +1,48 @@
1
+ from typing import Tuple, Union
2
+ from keras import ops
3
+ import numpy as np
4
+ from k3_node.ops.host import _in_shape_inference, to_numpy
5
+
6
+
7
+ def decimation_indices(
8
+ ptr,
9
+ decimation_factor: Union[int, float],
10
+ ) -> Tuple[any, any]:
11
+ r"""Gets indices which downsample each point cloud by a decimation factor.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import decimation_indices
17
+
18
+ ptr = np.array([0, 4, 10]) # two graphs with 4 and 6 nodes
19
+ index, new_ptr = decimation_indices(ptr, decimation_factor=2) # keep every 2nd node per graph
20
+ print(tuple(index.shape), tuple(new_ptr.shape)) # (5,) (3,)
21
+ ```
22
+ """
23
+ if decimation_factor < 1:
24
+ raise ValueError(
25
+ f"The argument `decimation_factor` should be higher than (or "
26
+ f"equal to) 1 for downsampling. (got {decimation_factor})"
27
+ )
28
+
29
+ ptr_np = to_numpy(ptr)
30
+ batch_size = len(ptr_np) - 1
31
+ count = ptr_np[1:] - ptr_np[:-1]
32
+ if _in_shape_inference(): # `ptr` is placeholder zeros: keep one (valid) node per graph
33
+ count = np.maximum(count, 1)
34
+ decim_count = np.maximum(count // int(decimation_factor), 1)
35
+
36
+ decim_indices_list = []
37
+ for i in range(batch_size):
38
+ perm = np.random.permutation(count[i])[:decim_count[i]]
39
+ decim_indices_list.append(ptr_np[i] + perm)
40
+
41
+ decim_indices = np.concatenate(decim_indices_list, axis=0)
42
+ decim_ptr = np.concatenate([[0], np.cumsum(decim_count)])
43
+
44
+ return (
45
+ ops.convert_to_tensor(decim_indices, dtype=ptr.dtype),
46
+ ops.convert_to_tensor(decim_ptr, dtype=ptr.dtype),
47
+ )
48
+