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,159 @@
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_max
7
+ from k3_node.ops.host import to_numpy
8
+
9
+
10
+ def _max_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
+ return segment_max(x, cluster, num_segments=size)
15
+
16
+
17
+ def max_pool_x(
18
+ cluster,
19
+ x,
20
+ batch,
21
+ batch_size: Optional[int] = None,
22
+ size: Optional[int] = None,
23
+ ) -> Tuple[any, Optional[any]]:
24
+ r"""Max-pools node features according to the clustering defined in `cluster`.
25
+
26
+ Example:
27
+ ```python
28
+ import numpy as np
29
+ from k3_node.layers import max_pool_x
30
+
31
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
32
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
33
+ cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
34
+
35
+ x_pool, batch_pool = max_pool_x(cluster, x, batch)
36
+ print(tuple(x_pool.shape)) # (5, 8)
37
+ ```
38
+ """
39
+ if size is not None:
40
+ if batch_size is None:
41
+ batch_size = int(to_numpy(batch).max()) + 1
42
+ return _max_pool_x(cluster, x, batch_size * size), None
43
+
44
+ cluster, perm = consecutive_cluster(cluster)
45
+ x = _max_pool_x(cluster, x)
46
+ batch = pool_batch(perm, batch)
47
+ return x, batch
48
+
49
+
50
+ def max_pool(
51
+ cluster,
52
+ data,
53
+ transform: Optional[Callable] = None,
54
+ edge_index: Optional[any] = None,
55
+ edge_attr: Optional[any] = None,
56
+ batch: Optional[any] = None,
57
+ pos: Optional[any] = None,
58
+ ):
59
+ r"""Pools and coarsens a graph given by `data` according to `cluster`.
60
+
61
+ Example:
62
+ ```python
63
+ import numpy as np
64
+ from k3_node.layers import max_pool
65
+
66
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
67
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
68
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
69
+ cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
70
+
71
+ x_pool, edge_index_pool, batch_pool = max_pool(cluster, x, edge_index, batch=batch)
72
+ print(tuple(x_pool.shape)) # (5, 8)
73
+ ```
74
+ """
75
+ cluster, perm = consecutive_cluster(cluster)
76
+
77
+ if hasattr(data, "x"):
78
+ data = as_mutable_graph(data)
79
+ x = getattr(data, "x", None)
80
+ if x is not None:
81
+ data.x = _max_pool_x(cluster, x)
82
+
83
+ edge_index = getattr(data, "edge_index", None)
84
+ edge_attr = getattr(data, "edge_attr", None)
85
+ if edge_index is not None:
86
+ data.edge_index, data.edge_attr = pool_edge(cluster, edge_index, edge_attr)
87
+
88
+ batch = getattr(data, "batch", None)
89
+ if batch is not None:
90
+ data.batch = pool_batch(perm, batch)
91
+
92
+ pos = getattr(data, "pos", None)
93
+ if pos is not None:
94
+ data.pos = pool_pos(cluster, pos)
95
+
96
+ if transform is not None:
97
+ data = transform(data)
98
+
99
+ return data
100
+
101
+ # Raw tensor mode
102
+ pooled_x = _max_pool_x(cluster, data)
103
+ pooled_edge_index, pooled_edge_attr = (None, None)
104
+ if edge_index is not None:
105
+ pooled_edge_index, pooled_edge_attr = pool_edge(cluster, edge_index, edge_attr)
106
+ pooled_batch = pool_batch(perm, batch) if batch is not None else None
107
+
108
+ return pooled_x, pooled_edge_index, pooled_batch
109
+
110
+
111
+ def max_pool_neighbor_x(
112
+ data,
113
+ edge_index=None,
114
+ flow: str = "source_to_target",
115
+ ):
116
+ r"""Max-pools neighboring node features.
117
+
118
+ Example:
119
+ ```python
120
+ import numpy as np
121
+ from k3_node.layers import max_pool_neighbor_x
122
+
123
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
124
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
125
+
126
+ out = max_pool_neighbor_x(x, edge_index=edge_index) # pool each node with its neighbors
127
+ print(tuple(out.shape)) # (10, 8)
128
+ ```
129
+ """
130
+ if hasattr(data, "x"):
131
+ x = data.x
132
+ edge_index = data.edge_index
133
+ is_data_obj = True
134
+ else:
135
+ x = data
136
+ is_data_obj = False
137
+ if edge_index is None:
138
+ raise ValueError("edge_index must be provided if data is a tensor.")
139
+
140
+ num_nodes = getattr(data, "num_nodes", None) if is_data_obj else ops.shape(x)[0]
141
+ if num_nodes is None:
142
+ num_nodes = ops.shape(x)[0]
143
+
144
+ # Add self-loops
145
+ loop_idx = ops.arange(num_nodes, dtype=edge_index.dtype)
146
+ loop_edge = ops.stack([loop_idx, loop_idx], axis=0)
147
+ full_edge_index = ops.concatenate([edge_index, loop_edge], axis=1)
148
+
149
+ row = full_edge_index[0]
150
+ col = full_edge_index[1]
151
+ row, col = (row, col) if flow == "source_to_target" else (col, row)
152
+
153
+ col = ops.cast(col, dtype="int32")
154
+ x_src = ops.take(x, row, axis=0)
155
+ out_x = segment_max(x_src, col, num_segments=num_nodes)
156
+ if is_data_obj:
157
+ data.x = out_x
158
+ return data
159
+ return out_x
@@ -0,0 +1,145 @@
1
+ from typing import Optional, Tuple
2
+ from k3_node.layers.aggr.base import to_dense_batch
3
+ from keras import initializers, layers, ops
4
+
5
+ EPS = 1e-15
6
+
7
+
8
+ class MemPooling(layers.Layer):
9
+ r"""Memory based pooling layer from `"Memory-Based Graph Networks"
10
+ <https://arxiv.org/abs/2002.09518>`_ paper.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import MemPooling
16
+
17
+ x = np.random.rand(2, 5, 8).astype("float32") # dense batch: 2 graphs with 5 nodes each
18
+ layer = MemPooling(in_channels=8, out_channels=16, heads=2, num_clusters=3)
19
+ x_pool, assignment = layer(x)
20
+ print(tuple(x_pool.shape)) # (2, 3, 16)
21
+ print(tuple(assignment.shape)) # (2, 5, 3): soft assignment of nodes to clusters
22
+ ```
23
+ """
24
+ def __init__(
25
+ self,
26
+ in_channels: int,
27
+ out_channels: int,
28
+ heads: int,
29
+ num_clusters: int,
30
+ tau: float = 1.0,
31
+ **kwargs,
32
+ ):
33
+ super().__init__(**kwargs)
34
+ self.in_channels = in_channels
35
+ self.out_channels = out_channels
36
+ self.heads = heads
37
+ self.num_clusters = num_clusters
38
+ self.tau = tau
39
+
40
+ self.k = self.add_weight(
41
+ shape=(heads, num_clusters, in_channels),
42
+ initializer=initializers.RandomUniform(minval=-1.0, maxval=1.0),
43
+ trainable=True,
44
+ name="k",
45
+ )
46
+ self.conv_weight = self.add_weight(
47
+ shape=(heads, 1),
48
+ initializer=initializers.GlorotUniform(),
49
+ trainable=True,
50
+ name="conv_weight",
51
+ )
52
+ self.lin = layers.Dense(out_channels, use_bias=False, name="lin")
53
+
54
+ def build(self, input_shape=None):
55
+ self.lin.build((None, self.num_clusters, self.in_channels))
56
+ super().build(input_shape)
57
+
58
+ def reset_parameters(self):
59
+ init = initializers.RandomUniform(minval=-1.0, maxval=1.0)
60
+ self.k.assign(init(self.k.shape, dtype=self.k.dtype))
61
+ glorot = initializers.GlorotUniform()
62
+ self.conv_weight.assign(glorot(self.conv_weight.shape, dtype=self.conv_weight.dtype))
63
+
64
+ @staticmethod
65
+ def kl_loss(S) -> any:
66
+ r"""The additional KL divergence-based loss."""
67
+ S_2 = ops.power(S, 2)
68
+ P = S_2 / ops.sum(S, axis=1, keepdims=True)
69
+ denom = ops.sum(P, axis=2, keepdims=True)
70
+ denom = ops.where(ops.sum(S, axis=2, keepdims=True) == 0.0, 1.0, denom)
71
+ P = P / denom
72
+
73
+ S_clamped = ops.maximum(S, EPS)
74
+ P_clamped = ops.maximum(P, EPS)
75
+ # KL(P || S) = sum(P * (log(P) - log(S)))
76
+ kl = P * (ops.log(P_clamped) - ops.log(S_clamped))
77
+ return ops.mean(ops.sum(kl, axis=(1, 2)))
78
+
79
+ def call(
80
+ self,
81
+ x,
82
+ batch: Optional[any] = None,
83
+ mask: Optional[any] = None,
84
+ max_num_nodes: Optional[int] = None,
85
+ batch_size: Optional[int] = None,
86
+ ) -> Tuple[any, any]:
87
+ r"""Forward pass."""
88
+ if len(ops.shape(x)) == 2:
89
+ # Node-level input: one dense [num_nodes, channels] block per graph (as in PyG)
90
+ if batch is None:
91
+ batch = ops.zeros((ops.shape(x)[0],), dtype="int32")
92
+ x, mask = to_dense_batch(x, ops.cast(batch, "int32"), dim_size=batch_size,
93
+ max_num_elements=max_num_nodes)
94
+ elif mask is None:
95
+ mask = ops.ones((ops.shape(x)[0], ops.shape(x)[1]), dtype=bool)
96
+
97
+ B = ops.shape(x)[0]
98
+ N = ops.shape(x)[1]
99
+ H = self.heads
100
+ K = self.num_clusters
101
+
102
+ # Compute pairwise squared Euclidean distance between k and x
103
+ # k: [H, K, C], x: [B, N, C]
104
+ # Reshape to [H * K, C] and [B * N, C]
105
+ k_flat = ops.reshape(self.k, (H * K, self.in_channels))
106
+ x_flat = ops.reshape(x, (B * N, self.in_channels))
107
+
108
+ k_sq = ops.sum(ops.power(k_flat, 2), axis=-1, keepdims=True) # [HK, 1]
109
+ x_sq = ops.sum(ops.power(x_flat, 2), axis=-1, keepdims=True) # [BN, 1]
110
+ dot = ops.matmul(k_flat, ops.transpose(x_flat)) # [HK, BN]
111
+ dist = ops.maximum(k_sq + ops.transpose(x_sq) - 2.0 * dot, 0.0) # [HK, BN]
112
+
113
+ dist = ops.power(1.0 + dist / self.tau, -(self.tau + 1.0) / 2.0)
114
+ # Reshape to [H, K, B, N] then permute to [B, H, N, K]
115
+ dist = ops.reshape(dist, (H, K, B, N))
116
+ dist = ops.transpose(dist, (2, 0, 3, 1))
117
+
118
+ S = dist / ops.sum(dist, axis=-1, keepdims=True) # [B, H, N, K]
119
+
120
+ # Conv over head dimension: reduce H -> 1
121
+ # conv_weight: [H, 1]
122
+ # Transpose S to [B, N, K, H] and multiply by [H, 1]
123
+ S_perm = ops.transpose(S, (0, 2, 3, 1)) # [B, N, K, H]
124
+ S_conv = ops.squeeze(ops.matmul(S_perm, self.conv_weight), axis=-1) # [B, N, K]
125
+ # [B, N, K]. With one cluster the softmax is 1 with zero gradient (as in PyG); `0 * S_conv`
126
+ # keeps the keys and conv weights in the graph, so Keras doesn't warn about them.
127
+ S = ops.softmax(S_conv, axis=-1) if K > 1 else 1.0 + 0.0 * S_conv
128
+
129
+ mask_f = ops.cast(ops.reshape(mask, (B, N, 1)), S.dtype)
130
+ S = S * mask_f
131
+
132
+ # x_out: [B, K, out_channels]
133
+ # S.transpose(1, 2) is [B, K, N]
134
+ # x is [B, N, C]
135
+ pooled_x = ops.matmul(ops.swapaxes(S, 1, 2), x) # [B, K, C]
136
+ x_out = self.lin(pooled_x)
137
+
138
+ return x_out, S
139
+
140
+ def __repr__(self) -> str:
141
+ return (
142
+ f"{self.__class__.__name__}({self.in_channels}, "
143
+ f"{self.out_channels}, heads={self.heads}, "
144
+ f"num_clusters={self.num_clusters})"
145
+ )
@@ -0,0 +1,144 @@
1
+ from typing import Callable, Optional, Tuple, Union
2
+ from keras import initializers, layers, ops
3
+
4
+ from .connect.filter_edges import FilterEdges
5
+ from .select.topk import SelectTopK
6
+ from k3_node.ops.segment import segment_sum
7
+ from k3_node.ops.creation import full
8
+
9
+
10
+ class PANPooling(layers.Layer):
11
+ r"""The path integral based pooling operator from the
12
+ `"Path Integral Based Convolution and Pooling for Graph Neural Networks"
13
+ <https://arxiv.org/abs/2006.16811>`_ paper.
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import PANPooling, PANConv
19
+
20
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
21
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
22
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
23
+
24
+ # PANPooling consumes the path weights produced by PANConv
25
+ x, weights = PANConv(in_channels=8, out_channels=8, filter_size=2)(x, edge_index)
26
+ layer = PANPooling(in_channels=8, ratio=0.5)
27
+ x_pool, edge_index_pool, edge_attr_pool, batch_pool, perm, score = layer(x, weights, batch=batch)
28
+ print(tuple(x_pool.shape)) # (6, 8)
29
+ ```
30
+ """
31
+ def __init__(
32
+ self,
33
+ in_channels: int,
34
+ ratio: float = 0.5,
35
+ min_score: Optional[float] = None,
36
+ multiplier: float = 1.0,
37
+ nonlinearity: Union[str, Callable] = "tanh",
38
+ **kwargs,
39
+ ):
40
+ super().__init__(**kwargs)
41
+
42
+ self.in_channels = in_channels
43
+ self.ratio = ratio
44
+ self.min_score = min_score
45
+ self.multiplier = multiplier
46
+ self.nonlinearity = nonlinearity
47
+
48
+ self.p = self.add_weight(
49
+ shape=(in_channels,),
50
+ initializer=initializers.Constant(1.0),
51
+ trainable=True,
52
+ name="p",
53
+ )
54
+ self.beta = self.add_weight(
55
+ shape=(2,),
56
+ initializer=initializers.Constant(0.5),
57
+ trainable=True,
58
+ name="beta",
59
+ )
60
+
61
+ self.select = SelectTopK(1, ratio, min_score, nonlinearity)
62
+ self.connect = FilterEdges()
63
+
64
+ def reset_parameters(self):
65
+ self.p.assign(ops.ones(self.p.shape, dtype=self.p.dtype))
66
+ self.beta.assign(full(self.beta.shape, 0.5, dtype=self.beta.dtype))
67
+ self.select.reset_parameters()
68
+
69
+ def build(self, input_shape=None):
70
+ if hasattr(self.select, "built") and not self.select.built:
71
+ self.select.build(None)
72
+ if hasattr(self.connect, "built") and not self.connect.built:
73
+ self.connect.build(None)
74
+ self.built = True
75
+
76
+ def call(
77
+ self,
78
+ x,
79
+ M,
80
+ batch: Optional[any] = None,
81
+ ) -> Tuple[any, any, any, Optional[any], any, any]:
82
+ r"""Forward pass.
83
+
84
+ Args:
85
+ x: Node feature matrix.
86
+ M: MET matrix, either a tuple/list (edge_index, edge_weight) or an object
87
+ with `.coo()` method (like PyG SparseTensor).
88
+ batch: Batch vector.
89
+ """
90
+ num_nodes = ops.shape(x)[0]
91
+ if batch is None:
92
+ batch = ops.zeros((num_nodes,), dtype="int32")
93
+
94
+ if hasattr(M, "coo"):
95
+ row, col, edge_weight = M.coo()
96
+ elif isinstance(M, (tuple, list)) and len(M) >= 2 and not isinstance(M[0], int):
97
+ edge_index, edge_weight = M[0], M[1]
98
+ row, col = edge_index[0], edge_index[1]
99
+ elif hasattr(M, "shape") and len(M.shape) == 2 and M.shape[0] == 2:
100
+ row, col = M[0], M[1]
101
+ edge_weight = ops.ones((ops.shape(col)[0],), dtype=x.dtype)
102
+ else:
103
+ nz = ops.where(M != 0)
104
+ row, col = nz[0], nz[1]
105
+ edge_weight = ops.take(ops.reshape(M, (-1,)), row * ops.shape(M)[1] + col, axis=0)
106
+
107
+ col = ops.cast(col, dtype="int32")
108
+ row = ops.cast(row, dtype="int32")
109
+ edge_weight = ops.cast(edge_weight, dtype=x.dtype)
110
+
111
+ score1 = ops.sum(x * self.p, axis=-1)
112
+ score2 = segment_sum(edge_weight, col, num_segments=num_nodes)
113
+ score = self.beta[0] * score1 + self.beta[1] * score2
114
+
115
+ select_out = self.select(score, batch)
116
+
117
+ perm = select_out.node_index
118
+ score_val = select_out.weight
119
+
120
+ x_pooled = ops.take(x, perm, axis=0) * ops.expand_dims(score_val, axis=-1)
121
+ if self.multiplier != 1.0:
122
+ x_pooled = x_pooled * self.multiplier
123
+
124
+ edge_index = ops.stack([col, row], axis=0)
125
+ connect_out = self.connect(select_out, edge_index, edge_weight, batch)
126
+
127
+ return (
128
+ x_pooled,
129
+ connect_out.edge_index,
130
+ connect_out.edge_attr,
131
+ connect_out.batch,
132
+ perm,
133
+ score_val,
134
+ )
135
+
136
+ def __repr__(self) -> str:
137
+ if self.min_score is None:
138
+ ratio = f"ratio={self.ratio}"
139
+ else:
140
+ ratio = f"min_score={self.min_score}"
141
+ return (
142
+ f"{self.__class__.__name__}({self.in_channels}, {ratio}, "
143
+ f"multiplier={self.multiplier})"
144
+ )
@@ -0,0 +1,212 @@
1
+ from typing import Optional
2
+ from keras import ops
3
+ import numpy as np
4
+
5
+ from .knn import knn
6
+ from k3_node.ops.host import to_numpy
7
+
8
+
9
+ def fps(
10
+ x,
11
+ batch: Optional[any] = None,
12
+ ratio: float = 0.5,
13
+ random_start: bool = True,
14
+ batch_size: Optional[int] = None,
15
+ ):
16
+ r"""Farthest Point Sampling algorithm.
17
+
18
+ Example:
19
+ ```python
20
+ import numpy as np
21
+ from k3_node.layers import fps
22
+
23
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
24
+ pos = np.random.rand(10, 3).astype("float32") # 3D positions
25
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
26
+
27
+ index = fps(pos, batch, ratio=0.4) # farthest point sampling: 40% of each graph's points
28
+ print(tuple(index.shape)) # (4,)
29
+ ```
30
+ """
31
+ x_np = to_numpy(x)
32
+ num_nodes = x_np.shape[0]
33
+
34
+ if batch is None:
35
+ batch_np = np.zeros(num_nodes, dtype=np.int64)
36
+ else:
37
+ batch_np = to_numpy(batch).astype(np.int64)
38
+
39
+ unique_batches = np.unique(batch_np)
40
+ selected_indices = []
41
+
42
+ for b in unique_batches:
43
+ idx_b = np.where(batch_np == b)[0]
44
+ n_b = len(idx_b)
45
+ if n_b == 0:
46
+ continue
47
+
48
+ if ratio >= 1:
49
+ num_samples = min(int(ratio), n_b)
50
+ else:
51
+ num_samples = max(1, int(np.ceil(ratio * n_b)))
52
+
53
+ pts_b = x_np[idx_b]
54
+ sampled = []
55
+
56
+ start_idx = np.random.randint(n_b) if random_start else 0
57
+ sampled.append(start_idx)
58
+ min_dist = np.sum((pts_b - pts_b[start_idx]) ** 2, axis=-1)
59
+
60
+ for _ in range(1, num_samples):
61
+ next_idx = int(np.argmax(min_dist))
62
+ sampled.append(next_idx)
63
+ dist_new = np.sum((pts_b - pts_b[next_idx]) ** 2, axis=-1)
64
+ min_dist = np.minimum(min_dist, dist_new)
65
+
66
+ selected_indices.extend(idx_b[sampled])
67
+
68
+ return ops.convert_to_tensor(np.array(selected_indices, dtype=np.int64), dtype="int64")
69
+
70
+
71
+ def radius(
72
+ x,
73
+ y,
74
+ r: float,
75
+ batch_x: Optional[any] = None,
76
+ batch_y: Optional[any] = None,
77
+ max_num_neighbors: int = 32,
78
+ num_workers: int = 1,
79
+ batch_size: Optional[int] = None,
80
+ ):
81
+ r"""Finds for each element in `y` all points in `x` within distance `r`.
82
+
83
+ Example:
84
+ ```python
85
+ import numpy as np
86
+ from k3_node.layers import radius
87
+
88
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
89
+ pos = np.random.rand(10, 3).astype("float32") # 3D positions
90
+ query = np.random.rand(4, 3).astype("float32") # 4 query points
91
+
92
+ assign = radius(pos, query, r=0.5) # all points within distance 0.5 of each query point
93
+ print(assign.shape[0]) # 2: rows: (query index, point index)
94
+ ```
95
+ """
96
+ x_np = to_numpy(x)
97
+ y_np = to_numpy(y)
98
+ if x_np.ndim == 1:
99
+ x_np = x_np[:, None]
100
+ if y_np.ndim == 1:
101
+ y_np = y_np[:, None]
102
+
103
+ N = x_np.shape[0]
104
+ M = y_np.shape[0]
105
+
106
+ if batch_x is None:
107
+ batch_x_np = np.zeros(N, dtype=np.int64)
108
+ else:
109
+ batch_x_np = to_numpy(batch_x).astype(np.int64)
110
+
111
+ if batch_y is None:
112
+ batch_y_np = np.zeros(M, dtype=np.int64)
113
+ else:
114
+ batch_y_np = to_numpy(batch_y).astype(np.int64)
115
+
116
+ rows = []
117
+ cols = []
118
+ r_sq = r * r
119
+
120
+ for i in range(M):
121
+ valid_b = batch_x_np == batch_y_np[i]
122
+ valid_idx = np.where(valid_b)[0]
123
+ if len(valid_idx) == 0:
124
+ continue
125
+
126
+ diff = x_np[valid_idx] - y_np[i]
127
+ dist_sq = np.sum(diff**2, axis=-1)
128
+ within_r = np.where(dist_sq <= r_sq)[0]
129
+ if len(within_r) > max_num_neighbors:
130
+ # Sort and take top max_num_neighbors
131
+ sort_order = np.argsort(dist_sq[within_r])[:max_num_neighbors]
132
+ within_r = within_r[sort_order]
133
+
134
+ for match_idx in valid_idx[within_r]:
135
+ rows.append(i)
136
+ cols.append(match_idx)
137
+
138
+ if len(rows) == 0:
139
+ return ops.zeros((2, 0), dtype="int64")
140
+
141
+ return ops.convert_to_tensor(np.stack([rows, cols], axis=0), dtype="int64")
142
+
143
+
144
+ def radius_graph(
145
+ x,
146
+ r: float,
147
+ batch: Optional[any] = None,
148
+ loop: bool = False,
149
+ max_num_neighbors: int = 32,
150
+ flow: str = "source_to_target",
151
+ num_workers: int = 1,
152
+ batch_size: Optional[int] = None,
153
+ ):
154
+ r"""Computes graph edges to all points within a given distance `r`.
155
+
156
+ Example:
157
+ ```python
158
+ import numpy as np
159
+ from k3_node.layers import radius_graph
160
+
161
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
162
+ pos = np.random.rand(10, 3).astype("float32") # 3D positions
163
+
164
+ edge_index = radius_graph(pos, r=0.5) # connect points closer than 0.5
165
+ print(edge_index.shape[0]) # 2
166
+ ```
167
+ """
168
+ assert flow in ["source_to_target", "target_to_source"]
169
+ edge_index = radius(
170
+ x,
171
+ x,
172
+ r,
173
+ batch_x=batch,
174
+ batch_y=batch,
175
+ max_num_neighbors=max_num_neighbors if loop else max_num_neighbors + 1,
176
+ )
177
+ edge_index_np = to_numpy(edge_index)
178
+
179
+ if not loop and edge_index_np.shape[1] > 0:
180
+ mask = edge_index_np[0] != edge_index_np[1]
181
+ edge_index_np = edge_index_np[:, mask]
182
+
183
+ if flow == "source_to_target" and edge_index_np.shape[1] > 0:
184
+ edge_index_np = np.flip(edge_index_np, axis=0)
185
+
186
+ return ops.convert_to_tensor(edge_index_np, dtype=edge_index.dtype)
187
+
188
+
189
+ def nearest(
190
+ x,
191
+ y,
192
+ batch_x: Optional[any] = None,
193
+ batch_y: Optional[any] = None,
194
+ ):
195
+ r"""Clusters each point in `x` to its nearest point in `y`.
196
+
197
+ Example:
198
+ ```python
199
+ import numpy as np
200
+ from k3_node.layers import nearest
201
+
202
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
203
+ pos = np.random.rand(10, 3).astype("float32") # 3D positions
204
+ query = np.random.rand(4, 3).astype("float32") # 4 query points
205
+
206
+ cluster = nearest(query, pos) # index of the closest point in `pos` for every query point
207
+ print(tuple(cluster.shape)) # (4,)
208
+ ```
209
+ """
210
+ edge_index = knn(y, x, k=1, batch_x=batch_y, batch_y=batch_x)
211
+ return edge_index[1]
212
+