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,110 @@
1
+ from typing import Callable, Union, Tuple
2
+ from keras import layers, ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class NNConv(MessagePassing):
8
+ r"""The continuous kernel-based convolutional operator from the
9
+ `"Neural Message Passing for Quantum Chemistry"
10
+ <https://arxiv.org/abs/1704.01212>`_ paper.
11
+
12
+ Args:
13
+ in_channels: Size of each input sample, or a tuple for bipartite graphs.
14
+ out_channels: Size of each output sample.
15
+ nn: A neural network :math:`h_{\mathbf{\Theta}}` that maps edge features
16
+ to shape :obj:`[-1, in_channels * out_channels]`.
17
+ aggr: The aggregation scheme to use (``"add"``, ``"mean"``, ``"max"``).
18
+ (default: ``"add"``)
19
+ root_weight: If set to :obj:`False`, the layer will not add the
20
+ transformed root node features. (default: ``True``)
21
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
22
+ (default: ``True``)
23
+
24
+ Example:
25
+ ```python
26
+ import numpy as np
27
+ import keras
28
+ from k3_node.layers import NNConv
29
+
30
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
31
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
32
+ edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
33
+
34
+ # `nn` maps each edge's features to an [in_channels * out_channels] weight matrix
35
+ layer = NNConv(in_channels=8, out_channels=16, nn=keras.layers.Dense(8 * 16))
36
+ out = layer(x, edge_index, edge_attr)
37
+ print(tuple(out.shape)) # (10, 16)
38
+ ```
39
+ """
40
+
41
+ def __init__(
42
+ self,
43
+ in_channels: Union[int, Tuple[int, int]],
44
+ out_channels: int,
45
+ nn: Callable,
46
+ aggr: str = "add",
47
+ root_weight: bool = True,
48
+ bias: bool = True,
49
+ **kwargs,
50
+ ):
51
+ super().__init__(aggr=aggr, **kwargs)
52
+ self.in_channels = in_channels
53
+ self.out_channels = out_channels
54
+ self.nn = nn
55
+ self.root_weight = root_weight
56
+ self.use_bias = bias
57
+
58
+ if isinstance(in_channels, int):
59
+ self.in_channels_src = in_channels
60
+ self.in_channels_dst = in_channels
61
+ else:
62
+ self.in_channels_src = in_channels[0]
63
+ self.in_channels_dst = in_channels[1]
64
+
65
+ if root_weight:
66
+ self.lin_root = layers.Dense(out_channels, use_bias=False)
67
+ else:
68
+ self.lin_root = None
69
+
70
+ def build(self, input_shape):
71
+ if self.lin_root is not None:
72
+ self.lin_root.build((None, self.in_channels_dst))
73
+
74
+ if self.use_bias:
75
+ self.bias = self.add_weight(
76
+ shape=(self.out_channels,),
77
+ initializer="zeros",
78
+ name="bias",
79
+ )
80
+ else:
81
+ self.bias = None
82
+ self.built = True
83
+
84
+ def call(self, x, edge_index=None, edge_attr=None, size=None, **kwargs):
85
+ if edge_index is None and isinstance(x, (tuple, list)):
86
+ x, edge_index = x[0], x[1]
87
+
88
+ if not isinstance(x, (tuple, list)):
89
+ x_src, x_dst = x, x
90
+ else:
91
+ x_src, x_dst = x[0], x[1]
92
+
93
+ weight = self.nn(edge_attr)
94
+ weight = ops.reshape(weight, (-1, self.in_channels_src, self.out_channels))
95
+
96
+ out = self.propagate(edge_index, x=x_src, weight=weight, size=size)
97
+
98
+ if self.root_weight and self.lin_root is not None and x_dst is not None:
99
+ out = out + self.lin_root(x_dst)
100
+
101
+ if self.bias is not None:
102
+ out = out + self.bias
103
+ return out
104
+
105
+ def message(self, x_j, weight):
106
+ # x_j: [E, in_channels_src], weight: [E, in_channels_src, out_channels]
107
+ x_j = ops.expand_dims(x_j, axis=1) # [E, 1, in_channels_src]
108
+ msg = ops.matmul(x_j, weight) # [E, 1, out_channels]
109
+ return ops.squeeze(msg, axis=1) # [E, out_channels]
110
+
@@ -0,0 +1,100 @@
1
+ from keras import ops
2
+ from keras.layers import Dense
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.ops.creation import scatter
6
+
7
+
8
+ class PANConv(MessagePassing):
9
+ r"""The path integral based convolution operator from the
10
+ `"Path Integral Based Convolution and Pooling for Graph Neural Networks"
11
+ <https://arxiv.org/abs/2004.14805>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import PANConv
17
+
18
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
19
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
20
+
21
+ layer = PANConv(in_channels=8, out_channels=16, filter_size=2)
22
+ out, weights = layer(x, edge_index) # also returns the learned path weights
23
+ print(tuple(out.shape)) # (10, 16)
24
+ ```
25
+ """
26
+ def __init__(
27
+ self,
28
+ in_channels: int,
29
+ out_channels: int,
30
+ filter_size: int,
31
+ **kwargs,
32
+ ):
33
+ kwargs.setdefault("aggr", "add")
34
+ super().__init__(**kwargs)
35
+
36
+ self.in_channels = in_channels
37
+ self.out_channels = out_channels
38
+ self.filter_size = filter_size
39
+
40
+ self.lin = Dense(out_channels, use_bias=True)
41
+ self.weight = self.add_weight(
42
+ shape=(filter_size + 1,),
43
+ initializer="ones",
44
+ name="weight",
45
+ )
46
+
47
+ def build(self, input_shape=None):
48
+ self.lin.build((None, self.in_channels))
49
+ self.built = True
50
+
51
+ def call(self, inputs, edge_index=None, **kwargs):
52
+ if edge_index is None:
53
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
54
+ x, edge_index = inputs
55
+ else:
56
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
57
+ else:
58
+ x = inputs
59
+
60
+ if not self.built:
61
+ self.build()
62
+
63
+ num_nodes = ops.shape(x)[0]
64
+
65
+ # Construct adjacency matrix
66
+ if hasattr(edge_index, "shape") and len(edge_index.shape) == 2 and edge_index.shape[0] == 2:
67
+ row, col = edge_index[0], edge_index[1]
68
+ adj = scatter(
69
+ ops.stack([row, col], axis=-1),
70
+ ops.ones((ops.shape(row)[0],), dtype=x.dtype),
71
+ shape=(num_nodes, num_nodes),
72
+ )
73
+ else:
74
+ adj = ops.cast(edge_index, x.dtype)
75
+
76
+ # PAN entropy / path calculation
77
+ # M = sum_{k=0}^filter_size weight[k] * A^k
78
+ # M = sum_{k=0}^filter_size exp(-E(k)/T) * A^k
79
+ w = ops.softplus(self.weight)
80
+ eye = ops.eye(num_nodes, dtype=x.dtype)
81
+ M = self.weight[0] * eye
82
+ M = w[0] * eye
83
+ curr_adj = eye
84
+ for k in range(1, self.filter_size + 1):
85
+ curr_adj = ops.matmul(curr_adj, adj)
86
+ M = M + self.weight[k] * curr_adj
87
+ M = M + w[k] * curr_adj
88
+
89
+ deg = ops.sum(M, axis=1)
90
+ deg_inv_sqrt = ops.power(ops.maximum(deg, 1e-12), -0.5)
91
+ # Avoid inf / nan
92
+ deg_inv_sqrt = ops.where(ops.isfinite(deg_inv_sqrt), deg_inv_sqrt, 0.0)
93
+
94
+ M_norm = ops.expand_dims(deg_inv_sqrt, 0) * M * ops.expand_dims(deg_inv_sqrt, 1)
95
+ M_norm = ops.expand_dims(deg_inv_sqrt, 1) * M * ops.expand_dims(deg_inv_sqrt, 0)
96
+
97
+ out = ops.matmul(M_norm, x)
98
+ out = self.lin(out)
99
+
100
+ return out, M_norm
@@ -0,0 +1,109 @@
1
+ from keras import ops
2
+ from keras.layers import Dense
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import gcn_norm
6
+
7
+
8
+ class PDNConv(MessagePassing):
9
+ r"""The pathfinder discovery network convolutional operator from the
10
+ `"Pathfinder Discovery Networks for Neural Message Passing"
11
+ <https://arxiv.org/abs/2010.12878>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import PDNConv
17
+
18
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
19
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
20
+ edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
21
+
22
+ layer = PDNConv(in_channels=8, out_channels=16, edge_dim=3, hidden_channels=16)
23
+ out = layer(x, edge_index, edge_attr)
24
+ print(tuple(out.shape)) # (10, 16)
25
+ ```
26
+ """
27
+ def __init__(
28
+ self,
29
+ in_channels: int,
30
+ out_channels: int,
31
+ edge_dim: int,
32
+ hidden_channels: int = 16,
33
+ add_self_loops: bool = True,
34
+ normalize: bool = True,
35
+ bias: bool = True,
36
+ **kwargs,
37
+ ):
38
+ kwargs.setdefault("aggr", "add")
39
+ super().__init__(**kwargs)
40
+
41
+ self.in_channels = in_channels
42
+ self.out_channels = out_channels
43
+ self.edge_dim = edge_dim
44
+ self.hidden_channels = hidden_channels
45
+ self.add_self_loops = add_self_loops
46
+ self.normalize = normalize
47
+ self.use_bias = bias
48
+
49
+ self.mlp_1 = Dense(hidden_channels, activation="relu")
50
+ self.mlp_2 = Dense(1, activation="sigmoid")
51
+ self.lin = Dense(out_channels, use_bias=False)
52
+
53
+ if bias:
54
+ self.bias = self.add_weight(
55
+ shape=(out_channels,),
56
+ initializer="zeros",
57
+ name="bias",
58
+ )
59
+ else:
60
+ self.bias = None
61
+
62
+ def build(self, input_shape=None):
63
+ self.mlp_1.build((None, self.edge_dim))
64
+ self.mlp_2.build((None, self.hidden_channels))
65
+ self.lin.build((None, self.in_channels))
66
+ self.built = True
67
+
68
+ def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
69
+ if edge_index is None:
70
+ if isinstance(inputs, (list, tuple)):
71
+ if len(inputs) == 3:
72
+ x, edge_index, edge_attr = inputs
73
+ elif len(inputs) == 2:
74
+ x, edge_index = inputs
75
+ else:
76
+ raise ValueError(f"Unexpected input length {len(inputs)}")
77
+ else:
78
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
79
+ else:
80
+ x = inputs
81
+
82
+ if not self.built:
83
+ self.build()
84
+
85
+ if edge_attr is not None:
86
+ edge_attr = self.mlp_1(edge_attr)
87
+ edge_attr = ops.squeeze(self.mlp_2(edge_attr), -1)
88
+
89
+ num_nodes = ops.shape(x)[0]
90
+ if self.normalize:
91
+ edge_index, edge_attr = gcn_norm(
92
+ edge_index,
93
+ edge_attr,
94
+ num_nodes=num_nodes,
95
+ add_self_loops=self.add_self_loops,
96
+ dtype=x.dtype,
97
+ )
98
+
99
+ x = self.lin(x)
100
+ out = self.propagate(edge_index, x=x, edge_weight=edge_attr)
101
+
102
+ if self.bias is not None:
103
+ out = out + self.bias
104
+
105
+ return out
106
+
107
+ def message(self, x_j, edge_weight=None):
108
+ return x_j if edge_weight is None else ops.expand_dims(edge_weight, -1) * x_j
109
+
@@ -0,0 +1,177 @@
1
+ from typing import Optional, Union, List, Callable
2
+ from keras import ops, activations
3
+ from keras.layers import Dense
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.aggr import DegreeScalerAggregation
7
+ from k3_node.ops.creation import repeat
8
+
9
+
10
+ class PNAConv(MessagePassing):
11
+ r"""The Principal Neighbourhood Aggregation graph convolutional operator
12
+ from the `"Principal Neighbourhood Aggregation for Graph Nets"
13
+ <https://arxiv.org/abs/2004.05718>`_ paper.
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import PNAConv
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
+
23
+ deg = np.array([0, 2, 4, 3, 1]) # in-degree histogram of the training graphs
24
+ layer = PNAConv(in_channels=8, out_channels=16, aggregators=["mean", "max", "min", "std"],
25
+ scalers=["identity", "amplification"], deg=deg)
26
+ out = layer(x, edge_index)
27
+ print(tuple(out.shape)) # (10, 16)
28
+ ```
29
+ """
30
+ @staticmethod
31
+ def get_degree_histogram(graphs):
32
+ r"""Returns the histogram of in-degrees over ``graphs`` (a dataset, a list of graphs or a
33
+ loader), which PNA uses to normalize its degree scalers:
34
+ ``deg[d]`` is the number of nodes with ``d`` incoming edges."""
35
+ import numpy as np
36
+
37
+ counts = []
38
+ for graph in graphs:
39
+ target = np.asarray(ops.convert_to_numpy(graph.edge_index))[1]
40
+ counts.append(np.bincount(np.bincount(target, minlength=graph.num_nodes)))
41
+ hist = np.zeros(max(len(c) for c in counts), dtype=np.int64)
42
+ for c in counts:
43
+ hist[: len(c)] += c
44
+ return hist
45
+
46
+ def __init__(
47
+ self,
48
+ in_channels: int,
49
+ out_channels: int,
50
+ aggregators: List[str],
51
+ scalers: List[str],
52
+ deg,
53
+ edge_dim: Optional[int] = None,
54
+ towers: int = 1,
55
+ pre_layers: int = 1,
56
+ post_layers: int = 1,
57
+ divide_input: bool = False,
58
+ act: Union[str, Callable, None] = "relu",
59
+ train_norm: bool = False,
60
+ **kwargs,
61
+ ):
62
+ aggr = DegreeScalerAggregation(aggregators, scalers, deg, train_norm)
63
+ super().__init__(aggr=aggr, node_dim=0, **kwargs)
64
+
65
+ if divide_input:
66
+ assert in_channels % towers == 0
67
+ assert out_channels % towers == 0
68
+
69
+ self.in_channels = in_channels
70
+ self.out_channels = out_channels
71
+ self.aggregators = aggregators
72
+ self.scalers = scalers
73
+ self.edge_dim = edge_dim
74
+ self.towers = towers
75
+ self.divide_input = divide_input
76
+ self.act = activations.get(act) if act is not None else None
77
+
78
+ self.F_in = in_channels // towers if divide_input else in_channels
79
+ self.F_out = out_channels // towers
80
+
81
+ if edge_dim is not None:
82
+ self.edge_encoder = Dense(self.F_in, use_bias=False)
83
+ else:
84
+ self.edge_encoder = None
85
+
86
+ self.pre_nns = []
87
+ self.post_nns = []
88
+ for _ in range(towers):
89
+ pre_dim = (3 if edge_dim else 2) * self.F_in
90
+ pre_layers_list = [Dense(self.F_in, use_bias=True)]
91
+ for _ in range(pre_layers - 1):
92
+ pre_layers_list.append(Dense(self.F_in, activation=self.act, use_bias=True))
93
+ self.pre_nns.append(pre_layers_list)
94
+
95
+ post_in_dim = (len(aggregators) * len(scalers) + 1) * self.F_in
96
+ post_layers_list = [Dense(self.F_out, use_bias=True)]
97
+ for _ in range(post_layers - 1):
98
+ post_layers_list.append(Dense(self.F_out, activation=self.act, use_bias=True))
99
+ self.post_nns.append(post_layers_list)
100
+
101
+ self.lin = Dense(out_channels, use_bias=True)
102
+
103
+ def build(self, input_shape=None):
104
+ if self.edge_encoder is not None:
105
+ self.edge_encoder.build((None, self.edge_dim))
106
+ for pre_list in self.pre_nns:
107
+ dim = (3 if self.edge_dim else 2) * self.F_in
108
+ for l in pre_list:
109
+ l.build((None, dim))
110
+ dim = self.F_in
111
+ for post_list in self.post_nns:
112
+ dim = (len(self.aggregators) * len(self.scalers) + 1) * self.F_in
113
+ for l in post_list:
114
+ l.build((None, dim))
115
+ dim = self.F_out
116
+ self.lin.build((None, self.out_channels))
117
+ self.built = True
118
+
119
+ def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
120
+ if edge_index is None:
121
+ if isinstance(inputs, (list, tuple)):
122
+ if len(inputs) == 3:
123
+ x, edge_index, edge_attr = inputs
124
+ elif len(inputs) == 2:
125
+ x, edge_index = inputs
126
+ else:
127
+ raise ValueError(f"Unexpected input length {len(inputs)}")
128
+ else:
129
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
130
+ else:
131
+ x = inputs
132
+
133
+ if not self.built:
134
+ self.build()
135
+
136
+ num_nodes = ops.shape(x)[0]
137
+ if self.divide_input:
138
+ x_towers = ops.reshape(x, (-1, self.towers, self.F_in))
139
+ else:
140
+ x_towers = repeat(ops.expand_dims(x, 1), self.towers, axis=1)
141
+
142
+ out = self.propagate(
143
+ edge_index,
144
+ x=x_towers,
145
+ edge_attr=edge_attr,
146
+ size=(num_nodes, num_nodes),
147
+ )
148
+
149
+ out = ops.concatenate([x_towers, out], axis=-1)
150
+
151
+ outs = []
152
+ for i, post_list in enumerate(self.post_nns):
153
+ h_i = out[:, i]
154
+ for l in post_list:
155
+ h_i = l(h_i)
156
+ outs.append(h_i)
157
+
158
+ out = ops.concatenate(outs, axis=-1)
159
+ return self.lin(out)
160
+
161
+ def message(self, x_i, x_j, edge_attr=None):
162
+ if edge_attr is not None and self.edge_encoder is not None:
163
+ edge_attr = self.edge_encoder(edge_attr)
164
+ edge_attr = repeat(ops.expand_dims(edge_attr, 1), self.towers, axis=1)
165
+ h = ops.concatenate([x_i, x_j, edge_attr], axis=-1)
166
+ else:
167
+ h = ops.concatenate([x_i, x_j], axis=-1)
168
+
169
+ hs = []
170
+ for i, pre_list in enumerate(self.pre_nns):
171
+ h_i = h[:, i]
172
+ for l in pre_list:
173
+ h_i = l(h_i)
174
+ hs.append(h_i)
175
+
176
+ return ops.stack(hs, axis=1)
177
+
@@ -0,0 +1,101 @@
1
+ from typing import Optional, Callable, Tuple
2
+ from keras import ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import remove_self_loops, add_self_loops
6
+
7
+
8
+ class PointNetConv(MessagePassing):
9
+ r"""The PointNet set abstraction layer from the `"PointNet++: Deep
10
+ Hierarchical Feature Learning on Point Sets in a Metric Space"
11
+ <https://arxiv.org/abs/1706.02413>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ import keras
17
+ from k3_node.layers import PointNetConv
18
+
19
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
20
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
21
+ pos = np.random.rand(10, 3).astype("float32") # 3D node positions
22
+
23
+ local_nn = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
24
+ layer = PointNetConv(local_nn=local_nn)
25
+ out = layer(x, pos, edge_index)
26
+ print(tuple(out.shape)) # (10, 16)
27
+ ```
28
+ """
29
+ def __init__(
30
+ self,
31
+ local_nn: Optional[Callable] = None,
32
+ global_nn: Optional[Callable] = None,
33
+ add_self_loops: bool = True,
34
+ **kwargs,
35
+ ):
36
+ kwargs.setdefault("aggr", "max")
37
+ super().__init__(**kwargs)
38
+
39
+ self.local_nn = local_nn
40
+ self.global_nn = global_nn
41
+ self.add_self_loops = add_self_loops
42
+
43
+ def build(self, input_shape=None):
44
+ self.built = True
45
+
46
+ def call(self, inputs, pos=None, edge_index=None, **kwargs):
47
+ if edge_index is None:
48
+ if isinstance(inputs, (list, tuple)):
49
+ if len(inputs) == 3:
50
+ x, pos, edge_index = inputs
51
+ elif len(inputs) == 2:
52
+ # x can be None or pos
53
+ x, edge_index = inputs
54
+ else:
55
+ raise ValueError(f"Unexpected input length {len(inputs)}")
56
+ else:
57
+ raise ValueError("Expected (x, pos, edge_index) or (x, edge_index)")
58
+ else:
59
+ x = inputs
60
+
61
+ if not self.built:
62
+ self.build()
63
+
64
+ if isinstance(pos, (list, tuple)):
65
+ pos_src, pos_dst = pos
66
+ else:
67
+ pos_src = pos_dst = pos
68
+
69
+ if isinstance(x, (list, tuple)):
70
+ x_src, x_dst = x
71
+ else:
72
+ x_src = x_dst = x
73
+
74
+ num_nodes = ops.shape(pos_dst)[0]
75
+ if self.add_self_loops:
76
+ edge_index, _ = remove_self_loops(edge_index)
77
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
78
+
79
+ out = self.propagate(
80
+ edge_index,
81
+ x=(x_src, x_dst),
82
+ pos=(pos_src, pos_dst),
83
+ size=(ops.shape(pos_src)[0], num_nodes),
84
+ )
85
+
86
+ if self.global_nn is not None:
87
+ out = self.global_nn(out)
88
+
89
+ return out
90
+
91
+ def message(self, x_j=None, pos_i=None, pos_j=None):
92
+ msg = pos_j - pos_i
93
+ if x_j is not None:
94
+ msg = ops.concatenate([x_j, msg], axis=-1)
95
+ if self.local_nn is not None:
96
+ msg = self.local_nn(msg)
97
+ return msg
98
+
99
+
100
+ PointConv = PointNetConv
101
+
@@ -0,0 +1,90 @@
1
+ from typing import Callable
2
+ from keras import ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class PointGNNConv(MessagePassing):
8
+ r"""The PointGNN graph convolutional operator from the
9
+ `"Point-GNN: Graph Neural Network for 3D Object Detection in a Point Cloud"
10
+ <https://arxiv.org/abs/2003.01251>`_ paper.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ import keras
16
+ from k3_node.layers import PointGNNConv
17
+
18
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
19
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
20
+ pos = np.random.rand(10, 3).astype("float32") # 3D node positions
21
+
22
+ layer = PointGNNConv(
23
+ mlp_h=keras.layers.Dense(3), # predicts a position offset per node
24
+ mlp_f=keras.layers.Dense(16), # edge feature network
25
+ mlp_g=keras.layers.Dense(8), # node update network (output size = input features)
26
+ )
27
+ out = layer(x, pos, edge_index)
28
+ print(tuple(out.shape)) # (10, 8)
29
+ ```
30
+ """
31
+ def __init__(
32
+ self,
33
+ mlp_h: Callable,
34
+ mlp_f: Callable,
35
+ mlp_g: Callable,
36
+ **kwargs,
37
+ ):
38
+ kwargs.setdefault("aggr", "max")
39
+ super().__init__(node_dim=0, **kwargs)
40
+
41
+ self.mlp_h = mlp_h
42
+ self.mlp_f = mlp_f
43
+ self.mlp_g = mlp_g
44
+
45
+ def build(self, input_shape=None):
46
+ self.built = True
47
+
48
+ def call(self, inputs, pos=None, edge_index=None, **kwargs):
49
+ if edge_index is None:
50
+ if isinstance(inputs, (list, tuple)):
51
+ if len(inputs) == 3:
52
+ x, pos, edge_index = inputs
53
+ elif len(inputs) == 2:
54
+ x, edge_index = inputs
55
+ else:
56
+ raise ValueError(f"Unexpected input length {len(inputs)}")
57
+ else:
58
+ raise ValueError("Expected (x, pos, edge_index)")
59
+ else:
60
+ x = inputs
61
+
62
+ if not self.built:
63
+ self.build()
64
+
65
+ if isinstance(x, (list, tuple)):
66
+ x_src, x_dst = x
67
+ else:
68
+ x_src = x_dst = x
69
+
70
+ if isinstance(pos, (list, tuple)):
71
+ pos_src, pos_dst = pos
72
+ else:
73
+ pos_src = pos_dst = pos
74
+
75
+ num_nodes = ops.shape(pos_dst)[0]
76
+ out = self.propagate(
77
+ edge_index,
78
+ x=(x_src, x_dst),
79
+ pos=(pos_src, pos_dst),
80
+ size=(ops.shape(pos_src)[0], num_nodes),
81
+ )
82
+
83
+ out = self.mlp_g(out)
84
+ return x_dst + out
85
+
86
+ def message(self, pos_j, pos_i, x_i, x_j):
87
+ delta = self.mlp_h(x_i)
88
+ e = ops.concatenate([pos_j - pos_i + delta, x_j], axis=-1)
89
+ return self.mlp_f(e)
90
+