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,456 @@
1
+ import keras
2
+ import pytest
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+ from k3_node.layers.pool import (
7
+ ASAPooling,
8
+ ApproxL2KNNIndex,
9
+ ApproxMIPSKNNIndex,
10
+ ClusterPooling,
11
+ EdgePooling,
12
+ FilterEdges,
13
+ KNNIndex,
14
+ L2KNNIndex,
15
+ LEConv,
16
+ MIPSKNNIndex,
17
+ MemPooling,
18
+ PANPooling,
19
+ SAGPooling,
20
+ SelectTopK,
21
+ TopKPooling,
22
+ approx_knn,
23
+ approx_knn_graph,
24
+ avg_pool,
25
+ avg_pool_neighbor_x,
26
+ avg_pool_x,
27
+ consecutive_cluster,
28
+ decimation_indices,
29
+ filter_adj,
30
+ fps,
31
+ global_add_pool,
32
+ global_max_pool,
33
+ global_mean_pool,
34
+ graclus,
35
+ knn,
36
+ knn_graph,
37
+ max_pool,
38
+ max_pool_neighbor_x,
39
+ max_pool_x,
40
+ nearest,
41
+ pool_batch,
42
+ pool_edge,
43
+ pool_pos,
44
+ radius,
45
+ radius_graph,
46
+ topk,
47
+ voxel_grid,
48
+ )
49
+ from k3_node.layers.pool.select.base import SelectOutput
50
+
51
+
52
+ def test_global_pooling():
53
+ x = ops.convert_to_tensor([
54
+ [1.0, 2.0],
55
+ [3.0, 4.0],
56
+ [5.0, 6.0],
57
+ [7.0, 8.0],
58
+ ], dtype="float32")
59
+ batch = ops.convert_to_tensor([0, 0, 1, 1], dtype="int32")
60
+
61
+ add = global_add_pool(x, batch)
62
+ assert np.allclose(ops.convert_to_numpy(add), [[4.0, 6.0], [12.0, 14.0]])
63
+
64
+ mean = global_mean_pool(x, batch)
65
+ assert np.allclose(ops.convert_to_numpy(mean), [[2.0, 3.0], [6.0, 7.0]])
66
+
67
+ max_p = global_max_pool(x, batch)
68
+ assert np.allclose(ops.convert_to_numpy(max_p), [[3.0, 4.0], [7.0, 8.0]])
69
+
70
+
71
+ def test_consecutive_cluster():
72
+ src = ops.convert_to_tensor([1, 4, 4, 7, 7, 7, 10], dtype="int32")
73
+ out, perm = consecutive_cluster(src)
74
+ assert np.array_equal(ops.convert_to_numpy(out), [0, 1, 1, 2, 2, 2, 3])
75
+ assert ops.shape(perm)[0] == 4
76
+
77
+
78
+ def test_filter_adj_and_layer():
79
+ edge_index = ops.convert_to_tensor([
80
+ [0, 1, 2, 3, 0],
81
+ [1, 2, 3, 0, 2],
82
+ ], dtype="int32")
83
+ edge_attr = ops.convert_to_tensor([1.0, 2.0, 3.0, 4.0, 5.0], dtype="float32")
84
+ node_index = ops.convert_to_tensor([0, 1, 2], dtype="int32")
85
+
86
+ new_edge_index, new_edge_attr = filter_adj(edge_index, edge_attr, node_index=node_index)
87
+ assert ops.shape(new_edge_index)[0] == 2
88
+ assert ops.shape(new_edge_index)[1] == 3 # (0,1), (1,2), (0,2)
89
+
90
+ # FilterEdges
91
+ layer = FilterEdges()
92
+ select_out = SelectOutput(
93
+ node_index=node_index,
94
+ num_nodes=4,
95
+ cluster_index=ops.arange(3, dtype="int32"),
96
+ num_clusters=3,
97
+ weight=ops.ones((3,), dtype="float32"),
98
+ )
99
+ batch = ops.zeros((4,), dtype="int32")
100
+ connect_out = layer(select_out, edge_index, edge_attr, batch)
101
+ assert ops.shape(connect_out.edge_index)[1] == 3
102
+ assert ops.shape(connect_out.batch)[0] == 3
103
+
104
+
105
+ def test_topk_and_select_topk():
106
+ x = ops.convert_to_tensor([2.0, 1.0, 5.0, 3.0, 8.0, 4.0], dtype="float32")
107
+ batch = ops.convert_to_tensor([0, 0, 0, 1, 1, 1], dtype="int32")
108
+
109
+ # Ratio test
110
+ perm = topk(x, ratio=0.5, batch=batch)
111
+ perm_np = ops.convert_to_numpy(perm)
112
+ assert 2 in perm_np[:2] # from first graph (5.0, 2.0)
113
+ assert 4 in perm_np[1:] # from second graph (8.0, 4.0)
114
+
115
+ # Min score test
116
+ perm_min = topk(x, ratio=None, batch=batch, min_score=3.5)
117
+ perm_min_np = ops.convert_to_numpy(perm_min)
118
+ assert 2 in perm_min_np # 5.0
119
+ assert 4 in perm_min_np # 8.0
120
+ assert 5 in perm_min_np # 4.0
121
+
122
+ # SelectTopK layer
123
+ layer = SelectTopK(in_channels=4, ratio=0.5)
124
+ feat = ops.ones((6, 4), dtype="float32")
125
+ out = layer(feat, batch)
126
+ assert isinstance(out, SelectOutput)
127
+ assert ops.shape(out.node_index)[0] == 4 # 2 per graph
128
+
129
+
130
+ def test_topk_pooling():
131
+ in_channels = 8
132
+ layer = TopKPooling(in_channels, ratio=0.5)
133
+
134
+ x = ops.ones((6, in_channels), dtype="float32")
135
+ edge_index = ops.convert_to_tensor([
136
+ [0, 1, 2, 3, 4, 5],
137
+ [1, 2, 0, 4, 5, 3],
138
+ ], dtype="int32")
139
+ batch = ops.convert_to_tensor([0, 0, 0, 1, 1, 1], dtype="int32")
140
+
141
+ out_x, out_edge_index, out_edge_attr, out_batch, perm, score = layer(
142
+ x, edge_index, batch=batch
143
+ )
144
+ assert ops.shape(out_x)[0] == 4 # 2 nodes per graph
145
+ assert ops.shape(out_x)[1] == in_channels
146
+ assert ops.shape(perm)[0] == 4
147
+ assert ops.shape(out_batch)[0] == 4
148
+
149
+
150
+ def test_sag_pooling():
151
+ in_channels = 8
152
+ layer = SAGPooling(in_channels, ratio=0.5)
153
+
154
+ x = ops.ones((4, in_channels), dtype="float32")
155
+ edge_index = ops.convert_to_tensor([
156
+ [0, 1, 2, 3, 0],
157
+ [1, 2, 3, 0, 2],
158
+ ], dtype="int32")
159
+
160
+ out_x, out_edge_index, out_edge_attr, out_batch, perm, score = layer(x, edge_index)
161
+ assert ops.shape(out_x)[0] == 2
162
+ assert ops.shape(out_x)[1] == in_channels
163
+ assert ops.shape(perm)[0] == 2
164
+
165
+
166
+ def test_edge_pooling():
167
+ in_channels = 8
168
+ layer = EdgePooling(in_channels)
169
+
170
+ x = ops.ones((4, in_channels), dtype="float32")
171
+ edge_index = ops.convert_to_tensor([
172
+ [0, 1, 2, 3],
173
+ [1, 2, 3, 0],
174
+ ], dtype="int32")
175
+ batch = ops.zeros((4,), dtype="int32")
176
+
177
+ new_x, new_edge_index, new_batch, unpool_info = layer(x, edge_index, batch)
178
+ assert ops.shape(new_x)[0] <= 4
179
+ assert ops.shape(new_x)[1] == in_channels
180
+
181
+ # Unpool
182
+ unpooled_x, _, _ = layer.unpool(new_x, unpool_info)
183
+ assert ops.shape(unpooled_x)[0] == 4
184
+ assert ops.shape(unpooled_x)[1] == in_channels
185
+
186
+
187
+ def test_cluster_pooling():
188
+ in_channels = 8
189
+ layer = ClusterPooling(in_channels)
190
+
191
+ x = ops.ones((4, in_channels), dtype="float32")
192
+ edge_index = ops.convert_to_tensor([
193
+ [0, 1, 2, 3],
194
+ [1, 2, 3, 0],
195
+ ], dtype="int32")
196
+ batch = ops.zeros((4,), dtype="int32")
197
+
198
+ new_x, new_edge_index, new_batch, unpool_info = layer(x, edge_index, batch)
199
+ assert unpool_info.edge_index is not None
200
+ assert unpool_info.cluster is not None
201
+ assert unpool_info.batch is not None
202
+
203
+
204
+ def test_mem_pooling():
205
+ in_channels, out_channels, heads, num_clusters = 8, 16, 2, 4
206
+ layer = MemPooling(in_channels, out_channels, heads=heads, num_clusters=num_clusters)
207
+
208
+ x = ops.ones((2, 5, in_channels), dtype="float32")
209
+ out_x, S = layer(x)
210
+ assert ops.shape(out_x) == (2, num_clusters, out_channels)
211
+ assert ops.shape(S) == (2, 5, num_clusters)
212
+
213
+ loss = MemPooling.kl_loss(S)
214
+ assert ops.shape(loss) == ()
215
+
216
+
217
+ def test_asap_pooling():
218
+ in_channels = 8
219
+ leconv = LEConv(in_channels, in_channels)
220
+ x = ops.ones((4, in_channels), dtype="float32")
221
+ edge_index = ops.convert_to_tensor([[0, 1, 2], [1, 2, 0]], dtype="int32")
222
+ le_out = leconv(x, edge_index)
223
+ assert ops.shape(le_out) == (4, in_channels)
224
+
225
+ asap = ASAPooling(in_channels, ratio=0.5)
226
+ out_x, out_edge_index, _, _, perm = asap(x, edge_index)
227
+ assert ops.shape(out_x)[0] == 2
228
+ assert ops.shape(out_x)[1] == in_channels
229
+
230
+
231
+ def test_pan_pooling():
232
+ in_channels = 8
233
+ pan = PANPooling(in_channels, ratio=0.5)
234
+ x = ops.ones((4, in_channels), dtype="float32")
235
+ M = ops.eye(4, dtype="float32")
236
+
237
+ out_x, out_edge_index, _, _, perm, score = pan(x, M)
238
+ assert ops.shape(out_x)[0] == 2
239
+ assert ops.shape(out_x)[1] == in_channels
240
+ assert ops.shape(out_edge_index)[0] == 2
241
+
242
+
243
+ def test_max_and_avg_pool():
244
+ cluster = ops.convert_to_tensor([0, 0, 1, 1], dtype="int32")
245
+ x = ops.convert_to_tensor([
246
+ [1.0, 5.0],
247
+ [3.0, 2.0],
248
+ [4.0, 8.0],
249
+ [6.0, 7.0],
250
+ ], dtype="float32")
251
+ edge_index = ops.convert_to_tensor([
252
+ [0, 1, 2, 3],
253
+ [1, 2, 3, 0],
254
+ ], dtype="int32")
255
+ batch = ops.zeros((4,), dtype="int32")
256
+
257
+ # max_pool_x & avg_pool_x
258
+ mx, mb = max_pool_x(cluster, x, batch)
259
+ assert np.allclose(ops.convert_to_numpy(mx), [[3.0, 5.0], [6.0, 8.0]])
260
+ ax, ab = avg_pool_x(cluster, x, batch)
261
+ assert np.allclose(ops.convert_to_numpy(ax), [[2.0, 3.5], [5.0, 7.5]])
262
+
263
+ # neighbor_x with raw tensors
264
+ mn_x = max_pool_neighbor_x(x, edge_index=edge_index)
265
+ assert ops.shape(mn_x) == (4, 2)
266
+ an_x = avg_pool_neighbor_x(x, edge_index=edge_index)
267
+ assert ops.shape(an_x) == (4, 2)
268
+
269
+ # max_pool & avg_pool (on graph)
270
+ m_x, m_edge, m_batch = max_pool(cluster, x, edge_index, batch=batch)
271
+ assert ops.shape(m_x)[0] == 2
272
+ a_x, a_edge, a_batch = avg_pool(cluster, x, edge_index, batch=batch)
273
+ assert ops.shape(a_x)[0] == 2
274
+
275
+
276
+ def test_pool_helpers():
277
+ cluster = ops.convert_to_tensor([0, 0, 1, 1], dtype="int32")
278
+ perm = ops.convert_to_tensor([0, 2], dtype="int32")
279
+ batch = ops.convert_to_tensor([0, 0, 0, 0], dtype="int32")
280
+ pos = ops.convert_to_tensor([[0.0, 0.0], [1.0, 1.0], [2.0, 2.0], [3.0, 3.0]], dtype="float32")
281
+ edge_index = ops.convert_to_tensor([[0, 1, 2], [1, 2, 3]], dtype="int32")
282
+ edge_attr = ops.convert_to_tensor([[1.0], [2.0], [3.0]], dtype="float32")
283
+
284
+ p_edge, p_attr = pool_edge(cluster, edge_index, edge_attr)
285
+ assert ops.shape(p_edge)[0] == 2
286
+ assert p_attr is not None
287
+
288
+ p_b = pool_batch(perm, batch)
289
+ assert ops.shape(p_b)[0] == 2
290
+
291
+ p_p = pool_pos(cluster, pos)
292
+ assert ops.shape(p_p)[0] == 2
293
+
294
+
295
+ def test_voxel_grid():
296
+ pos = ops.convert_to_tensor([
297
+ [0.0, 0.0],
298
+ [0.5, 0.5],
299
+ [1.2, 1.2],
300
+ [2.1, 2.1],
301
+ ], dtype="float32")
302
+ size = ops.convert_to_tensor([1.0, 1.0], dtype="float32")
303
+
304
+ c = voxel_grid(pos, size)
305
+ c_np = ops.convert_to_numpy(c)
306
+ assert c_np[0] == c_np[1]
307
+ assert c_np[0] != c_np[2]
308
+ assert c_np[2] != c_np[3]
309
+
310
+
311
+ def test_graclus():
312
+ edge_index = ops.convert_to_tensor([
313
+ [0, 1, 2, 3],
314
+ [1, 2, 3, 0],
315
+ ], dtype="int32")
316
+ c = graclus(edge_index, num_nodes=4)
317
+ c_np = ops.convert_to_numpy(c)
318
+ assert len(np.unique(c_np)) == 2
319
+
320
+
321
+ def test_decimation():
322
+ ptr = ops.convert_to_tensor([0, 4, 10], dtype="int32")
323
+ dec_idx, dec_ptr = decimation_indices(ptr, decimation_factor=2)
324
+ assert ops.shape(dec_idx)[0] == 5 # 2 from first, 3 from second
325
+ assert ops.shape(dec_ptr)[0] == 3
326
+
327
+
328
+ def test_knn_and_indices():
329
+ x = ops.convert_to_tensor([
330
+ [0.0, 0.0],
331
+ [0.1, 0.1],
332
+ [1.0, 1.0],
333
+ [1.1, 1.1],
334
+ ], dtype="float32")
335
+ y = ops.convert_to_tensor([
336
+ [0.0, 0.0],
337
+ [1.0, 1.0],
338
+ ], dtype="float32")
339
+
340
+ # KNNIndex / L2KNNIndex
341
+ idx_layer = L2KNNIndex(x)
342
+ out = idx_layer.search(y, k=2)
343
+ assert ops.shape(out.score) == (2, 2)
344
+ assert ops.shape(out.index) == (2, 2)
345
+
346
+ # MIPSKNNIndex
347
+ mips_layer = MIPSKNNIndex(x)
348
+ mips_out = mips_layer.search(y, k=2)
349
+ assert ops.shape(mips_out.score) == (2, 2)
350
+ assert ops.shape(mips_out.index) == (2, 2)
351
+
352
+ # knn function
353
+ res = knn(x, y, k=2)
354
+ assert ops.shape(res) == (2, 4)
355
+
356
+ # knn_graph
357
+ g = knn_graph(x, k=2)
358
+ assert ops.shape(g)[0] == 2
359
+
360
+ # approx knn
361
+ approx_res = approx_knn(x, y, k=2)
362
+ assert ops.shape(approx_res) == (2, 4)
363
+ approx_g = approx_knn_graph(x, k=2)
364
+ assert ops.shape(approx_g)[0] == 2
365
+
366
+
367
+ def test_point_cloud():
368
+ pos = ops.convert_to_tensor([
369
+ [0.0, 0.0],
370
+ [0.1, 0.0],
371
+ [10.0, 10.0],
372
+ [10.1, 10.0],
373
+ ], dtype="float32")
374
+
375
+ # fps
376
+ fps_idx = fps(pos, ratio=0.5)
377
+ assert ops.shape(fps_idx)[0] == 2
378
+
379
+ # radius
380
+ rad_edge = radius(pos, pos, r=1.0)
381
+ assert ops.shape(rad_edge)[0] == 2
382
+ assert ops.shape(rad_edge)[1] >= 4 # at least self loops + close pairs
383
+
384
+ # radius_graph
385
+ rad_g = radius_graph(pos, r=1.0)
386
+ assert ops.shape(rad_g)[0] == 2
387
+
388
+ # nearest
389
+ near = nearest(pos, pos)
390
+ assert ops.shape(near)[0] == 4
391
+
392
+
393
+ def test_global_pool_with_unsorted_batch():
394
+ # Pooling layers (e.g. EdgePooling) can return batch vectors that are not sorted by graph.
395
+ import numpy as np
396
+ from keras import ops
397
+ from k3_node.layers import global_add_pool, global_mean_pool
398
+
399
+ x = np.arange(8, dtype="float32").reshape(4, 2)
400
+ batch = np.array([1, 0, 1, 0], dtype="int32")
401
+ np.testing.assert_allclose(ops.convert_to_numpy(global_add_pool(x, batch)), [[8, 10], [4, 6]])
402
+ np.testing.assert_allclose(ops.convert_to_numpy(global_mean_pool(x, batch)), [[4, 5], [2, 3]])
403
+
404
+
405
+ @pytest.mark.skipif(keras.backend.backend() != "jax", reason="jax.jit specific")
406
+ def test_global_pool_under_jax_jit_requires_size():
407
+ import jax
408
+ import numpy as np
409
+ from k3_node.layers import global_add_pool
410
+
411
+ x = np.arange(8, dtype="float32").reshape(4, 2)
412
+ batch = np.array([0, 0, 1, 1], dtype="int32")
413
+ with pytest.raises(ValueError, match="size="):
414
+ jax.jit(lambda x, b: global_add_pool(x, b))(x, batch)
415
+ out = jax.jit(lambda x, b: global_add_pool(x, b, size=2))(x, batch)
416
+ np.testing.assert_allclose(np.asarray(out), [[2, 4], [10, 12]])
417
+
418
+
419
+ @pytest.mark.skipif(keras.backend.backend() not in ("tensorflow", "jax"), reason="needs a compiling backend")
420
+ @pytest.mark.parametrize("layer_name", ["EdgePooling", "ClusterPooling"])
421
+ def test_host_side_pooling_refuses_compiled_execution(layer_name):
422
+ # These layers pick clusters on the host; inside tf.function / jax.jit they used to return
423
+ # their input unpooled without any error.
424
+ import numpy as np
425
+ import k3_node.layers as L
426
+
427
+ layer = getattr(L, layer_name)(4)
428
+ x = np.random.randn(6, 4).astype("float32")
429
+ edge_index = np.array([[0, 1, 2, 3, 4], [1, 2, 3, 4, 5]], dtype="int32")
430
+ batch = np.zeros(6, dtype="int32")
431
+ layer(x, edge_index, batch) # eager works
432
+
433
+ if keras.backend.backend() == "jax":
434
+ import jax
435
+
436
+ compiled = jax.jit(lambda x: layer(x, edge_index, batch)[0])
437
+ else:
438
+ import tensorflow as tf
439
+
440
+ compiled = tf.function(lambda x: layer(x, edge_index, batch)[0])
441
+ with pytest.raises(Exception, match="run_eagerly"):
442
+ compiled(x)
443
+
444
+
445
+ def test_mem_pooling_respects_batch():
446
+ import numpy as np
447
+ from k3_node.layers import MemPooling
448
+
449
+ x = np.random.rand(7, 4).astype("float32")
450
+ batch = np.array([0, 0, 0, 1, 1, 2, 2]) # three graphs
451
+ pool = MemPooling(4, 8, heads=2, num_clusters=3)
452
+ out, S = pool(x, batch)
453
+ assert tuple(out.shape) == (3, 3, 8) # one pooled set per graph
454
+ # Pooling a graph alone gives the same result as pooling it within the batch
455
+ alone, _ = pool(x[3:5], np.zeros(2, dtype="int32"))
456
+ np.testing.assert_allclose(ops.convert_to_numpy(out)[1], ops.convert_to_numpy(alone)[0], rtol=1e-5, atol=1e-6)
@@ -0,0 +1,103 @@
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
+
7
+
8
+ class TopKPooling(layers.Layer):
9
+ r""":math:`\mathrm{top}_k` pooling operator from the `"Graph U-Nets"
10
+ <https://arxiv.org/abs/1905.05178>`_, `"Towards Sparse
11
+ Hierarchical Graph Classifiers" <https://arxiv.org/abs/1811.01287>`_
12
+ and `"Understanding Attention and Generalization in Graph Neural
13
+ Networks" <https://arxiv.org/abs/1905.02850>`_ papers.
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import TopKPooling
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
+ layer = TopKPooling(in_channels=8, ratio=0.5) # keep half of the nodes of each graph
25
+ out = layer(x, edge_index, batch=batch)
26
+ x_pool, edge_index_pool, edge_attr_pool, batch_pool = out[0], out[1], out[2], out[3]
27
+ print(tuple(x_pool.shape)) # (6, 8)
28
+ ```
29
+ """
30
+ def __init__(
31
+ self,
32
+ in_channels: int,
33
+ ratio: Union[int, float] = 0.5,
34
+ min_score: Optional[float] = None,
35
+ multiplier: float = 1.0,
36
+ nonlinearity: Union[str, Callable] = "tanh",
37
+ **kwargs,
38
+ ):
39
+ super().__init__(**kwargs)
40
+
41
+ self.in_channels = in_channels
42
+ self.ratio = ratio
43
+ self.min_score = min_score
44
+ self.multiplier = multiplier
45
+ self.nonlinearity = nonlinearity
46
+
47
+ self.select = SelectTopK(in_channels, ratio, min_score, nonlinearity)
48
+ self.connect = FilterEdges()
49
+
50
+ def reset_parameters(self):
51
+ r"""Resets all learnable parameters of the module."""
52
+ self.select.reset_parameters()
53
+
54
+ def build(self, input_shape=None):
55
+ if not self.select.built:
56
+ self.select.build(input_shape)
57
+ if hasattr(self.connect, "built") and not self.connect.built:
58
+ self.connect.build(None)
59
+ self.built = True
60
+
61
+ def call(
62
+ self,
63
+ x,
64
+ edge_index,
65
+ edge_attr: Optional[any] = None,
66
+ batch: Optional[any] = None,
67
+ attn: Optional[any] = None,
68
+ ) -> Tuple[any, any, Optional[any], Optional[any], any, any]:
69
+ r"""Forward pass."""
70
+ if batch is None:
71
+ batch = ops.zeros((ops.shape(x)[0],), dtype="int32")
72
+
73
+ attn_input = x if attn is None else attn
74
+ select_out = self.select(attn_input, batch)
75
+
76
+ perm = select_out.node_index
77
+ score = select_out.weight
78
+
79
+ x_pooled = ops.take(x, perm, axis=0) * ops.expand_dims(score, axis=-1)
80
+ if self.multiplier != 1.0:
81
+ x_pooled = x_pooled * self.multiplier
82
+
83
+ connect_out = self.connect.call(select_out, edge_index, edge_attr, batch)
84
+
85
+ return (
86
+ x_pooled,
87
+ connect_out.edge_index,
88
+ connect_out.edge_attr,
89
+ connect_out.batch,
90
+ perm,
91
+ score,
92
+ )
93
+
94
+ def __repr__(self) -> str:
95
+ if self.min_score is None:
96
+ ratio = f"ratio={self.ratio}"
97
+ else:
98
+ ratio = f"min_score={self.min_score}"
99
+ return (
100
+ f"{self.__class__.__name__}({self.in_channels}, {ratio}, "
101
+ f"multiplier={self.multiplier})"
102
+ )
103
+
@@ -0,0 +1,70 @@
1
+ from typing import List, Optional, Union
2
+ from keras import ops
3
+ import numpy as np
4
+ from k3_node.ops.host import to_numpy
5
+
6
+
7
+ def voxel_grid(
8
+ pos,
9
+ size: Union[float, List[float], any],
10
+ batch: Optional[any] = None,
11
+ start: Optional[Union[float, List[float], any]] = None,
12
+ end: Optional[Union[float, List[float], any]] = None,
13
+ ):
14
+ r"""Voxel grid pooling that clusters points within the same voxel.
15
+
16
+ Example:
17
+ ```python
18
+ import numpy as np
19
+ from k3_node.layers import voxel_grid
20
+
21
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
22
+ pos = np.random.rand(10, 3).astype("float32") # 3D positions
23
+
24
+ cluster = voxel_grid(pos, size=0.5) # voxel id of each point
25
+ print(tuple(cluster.shape)) # (10,)
26
+ ```
27
+ """
28
+ pos_np = to_numpy(pos)
29
+ if pos_np.ndim == 1:
30
+ pos_np = pos_np[:, None]
31
+ dim = pos_np.shape[1]
32
+
33
+ if isinstance(size, (int, float)):
34
+ size_np = np.full(dim, float(size))
35
+ else:
36
+ size_np = np.array(size, dtype=np.float64)
37
+
38
+ if start is None:
39
+ start_np = np.min(pos_np, axis=0)
40
+ elif isinstance(start, (int, float)):
41
+ start_np = np.full(dim, float(start))
42
+ else:
43
+ start_np = np.array(start, dtype=np.float64)
44
+
45
+ if end is None:
46
+ end_np = np.max(pos_np, axis=0)
47
+ elif isinstance(end, (int, float)):
48
+ end_np = np.full(dim, float(end))
49
+ else:
50
+ end_np = np.array(end, dtype=np.float64)
51
+
52
+ # Number of bins per dimension
53
+ num_bins = np.floor((end_np - start_np) / size_np).astype(np.int64) + 1
54
+
55
+ # Discretized coordinates
56
+ c = np.floor((pos_np - start_np) / size_np).astype(np.int64)
57
+
58
+ # Compute flat index
59
+ cluster_np = np.zeros(len(pos_np), dtype=np.int64)
60
+ multiplier = 1
61
+ for d in range(dim):
62
+ cluster_np += c[:, d] * multiplier
63
+ multiplier *= num_bins[d]
64
+
65
+ if batch is not None:
66
+ batch_np = to_numpy(batch).astype(np.int64)
67
+ cluster_np += batch_np * multiplier
68
+
69
+ return ops.convert_to_tensor(cluster_np, dtype="int64")
70
+
@@ -0,0 +1,9 @@
1
+ r"""Unpooling package."""
2
+
3
+ from .knn_interpolate import knn_interpolate
4
+
5
+ __all__ = [
6
+ "knn_interpolate",
7
+ ]
8
+
9
+ classes = __all__
@@ -0,0 +1,57 @@
1
+ from keras import ops
2
+
3
+ from k3_node.layers.conv.utils import scatter
4
+ from k3_node.layers.pool.knn import knn
5
+
6
+
7
+ def knn_interpolate(x, pos_x, pos_y, batch_x=None, batch_y=None, k: int = 3, num_workers: int = 1):
8
+ r"""The k-NN interpolation from the `"PointNet++: Deep Hierarchical
9
+ Feature Learning on Point Sets in a Metric Space"
10
+ <https://arxiv.org/abs/1706.02413>`_ paper.
11
+
12
+ For each point :math:`y` with position :math:`\mathbf{p}(y)`, its
13
+ interpolated features :math:`\mathbf{f}(y)` are given by
14
+
15
+ .. math::
16
+ \mathbf{f}(y) = \frac{\sum_{i=1}^k w(x_i) \mathbf{f}(x_i)}{\sum_{i=1}^k
17
+ w(x_i)} \textrm{, where } w(x_i) = \frac{1}{d(\mathbf{p}(y),
18
+ \mathbf{p}(x_i))^2}
19
+
20
+ and :math:`\{ x_1, \ldots, x_k \}` denoting the :math:`k` nearest points
21
+ to :math:`y`.
22
+
23
+ Args:
24
+ x: Node feature matrix :math:`\mathbf{X} \in \mathbb{R}^{N \times F}`.
25
+ pos_x: Node position matrix :math:`\in \mathbb{R}^{N \times d}`.
26
+ pos_y: Upsampled node position matrix :math:`\in \mathbb{R}^{M \times d}`.
27
+ batch_x: Batch vector assigning each node from :math:`\mathbf{X}` to
28
+ a specific example. (default: :obj:`None`)
29
+ batch_y: Batch vector assigning each node from :math:`\mathbf{Y}` to
30
+ a specific example. (default: :obj:`None`)
31
+ k (int, optional): Number of neighbors. (default: :obj:`3`)
32
+ num_workers (int, optional): Unused, kept for API compatibility.
33
+
34
+ Example:
35
+ ```python
36
+ import numpy as np
37
+ from k3_node.layers import knn_interpolate
38
+
39
+ x = np.random.rand(6, 8).astype("float32") # features of 6 coarse points
40
+ pos_x = np.random.rand(6, 3).astype("float32")
41
+ pos_y = np.random.rand(10, 3).astype("float32") # 10 fine points to interpolate onto
42
+ out = knn_interpolate(x, pos_x, pos_y, k=3)
43
+ print(tuple(out.shape)) # (10, 8)
44
+ ```
45
+ """
46
+ assign_index = knn(pos_x, pos_y, k, batch_x=batch_x, batch_y=batch_y, num_workers=num_workers)
47
+ y_idx, x_idx = assign_index[0], assign_index[1]
48
+
49
+ diff = ops.take(pos_x, x_idx, axis=0) - ops.take(pos_y, y_idx, axis=0)
50
+ squared_distance = ops.sum(diff * diff, axis=-1, keepdims=True)
51
+ weights = 1.0 / ops.maximum(squared_distance, 1e-16)
52
+
53
+ num_y = ops.shape(pos_y)[0]
54
+ y = scatter(ops.take(x, x_idx, axis=0) * weights, y_idx, dim=0, dim_size=num_y, reduce="sum")
55
+ y = y / scatter(weights, y_idx, dim=0, dim_size=num_y, reduce="sum")
56
+
57
+ return y