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,153 @@
1
+ from typing import Optional
2
+ import numpy as np
3
+ import keras
4
+ from keras import ops
5
+ from k3_node.ops.segment import segment_sum
6
+
7
+
8
+ def to_dense_batch(x, batch=None, batch_size=None):
9
+ """Differentiable dense batching; see :func:`k3_node.layers.aggr.to_dense_batch`."""
10
+ from k3_node.layers.aggr.base import to_dense_batch as _to_dense_batch
11
+
12
+ if batch is None:
13
+ return ops.expand_dims(x, axis=0), ops.ones((1, ops.shape(x)[0]), dtype="bool")
14
+ return _to_dense_batch(x, batch, dim_size=batch_size)
15
+
16
+
17
+ class GPSConv(keras.layers.Layer):
18
+ r"""The general, powerful, scalable (GPS) graph transformer layer from the
19
+ `"Recipe for a General, Powerful, Scalable Graph Transformer"
20
+ <https://arxiv.org/abs/2205.12454>`_ paper.
21
+
22
+ Args:
23
+ channels (int): Size of each input sample.
24
+ conv (keras.layers.Layer, optional): The local message passing layer.
25
+ heads (int, optional): Number of multi-head-attentions. (default: :obj:`1`)
26
+ dropout (float, optional): Dropout probability. (default: :obj:`0.0`)
27
+ act (str, optional): Activation function. (default: :obj:`"relu"`)
28
+ norm (str, optional): Normalization function. (default: :obj:`"batch_norm"`)
29
+
30
+ Call arguments: ``x``, ``edge_index``, ``batch`` (the graph of every node) and ``batch_size``
31
+ (the number of graphs; pass it when training compiled so that shapes are static), plus any
32
+ arguments of the local ``conv`` such as ``edge_attr``.
33
+
34
+ Example:
35
+ ```python
36
+ import numpy as np
37
+ from k3_node.layers import GPSConv, GCNConv
38
+
39
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
40
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
41
+
42
+ batch = np.repeat([0, 1], 5) # two graphs with 5 nodes each
43
+ layer = GPSConv(channels=8, conv=GCNConv(8, 8), heads=2)
44
+ out = layer(x, edge_index, batch=batch)
45
+ print(tuple(out.shape)) # (10, 8)
46
+ ```
47
+ """
48
+
49
+ def __init__(
50
+ self,
51
+ channels: int,
52
+ conv: Optional[keras.layers.Layer] = None,
53
+ heads: int = 1,
54
+ dropout: float = 0.0,
55
+ act: str = "relu",
56
+ norm: Optional[str] = "batch_norm",
57
+ **kwargs,
58
+ ):
59
+ conv = kwargs.pop("local_gnn", conv)
60
+ super().__init__(**kwargs)
61
+
62
+ self.channels = channels
63
+ self.conv = conv
64
+ self.heads = heads
65
+ self.dropout_rate = dropout
66
+ self.act = act
67
+ self.norm_name = norm
68
+
69
+ self.attn = keras.layers.MultiHeadAttention(
70
+ num_heads=heads,
71
+ key_dim=channels // heads,
72
+ value_dim=channels // heads,
73
+ output_shape=channels,
74
+ )
75
+
76
+ self.mlp_l1 = keras.layers.Dense(channels * 2)
77
+ self.mlp_l2 = keras.layers.Dense(channels)
78
+ self.dropout = keras.layers.Dropout(dropout)
79
+
80
+ if norm == "batch_norm":
81
+ self.norm1 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5) if conv is not None else None
82
+ self.norm2 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
83
+ self.norm3 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
84
+ elif norm == "layer_norm":
85
+ self.norm1 = keras.layers.LayerNormalization(axis=-1) if conv is not None else None
86
+ self.norm2 = keras.layers.LayerNormalization(axis=-1)
87
+ self.norm3 = keras.layers.LayerNormalization(axis=-1)
88
+ else:
89
+ self.norm1 = None
90
+ self.norm2 = None
91
+ self.norm3 = None
92
+
93
+ def build(self, input_shape=None):
94
+ if self.conv is not None and hasattr(self.conv, "build") and not self.conv.built:
95
+ self.conv.build(input_shape)
96
+ if not self.attn.built:
97
+ self.attn.build((None, None, self.channels), (None, None, self.channels))
98
+ if not self.mlp_l1.built:
99
+ self.mlp_l1.build((None, self.channels))
100
+ if not self.mlp_l2.built:
101
+ self.mlp_l2.build((None, self.channels * 2))
102
+ super().build(input_shape)
103
+
104
+ def call(self, x, edge_index, batch=None, batch_size=None, training=None, **kwargs):
105
+ if not self.built:
106
+ self.build((None, self.channels))
107
+
108
+ # `training` is forwarded explicitly: Keras does not propagate it to nested layers on JAX.
109
+ hs = []
110
+ if self.conv is not None:
111
+ h = self.conv(x, edge_index, training=training, **kwargs)
112
+ h = self.dropout(h, training=training)
113
+ h = h + x
114
+ if self.norm1 is not None:
115
+ h = self.norm1(h, training=training)
116
+ hs.append(h)
117
+
118
+ # Global attention
119
+ # `batch_size` (the number of graphs) keeps shapes static when compiled
120
+ h_dense, mask = to_dense_batch(x, batch, batch_size)
121
+ # Attention mask for Keras: shape (B, 1, max_nodes)
122
+ attn_mask = ops.expand_dims(mask, axis=1)
123
+ attn_out = self.attn(h_dense, h_dense, attention_mask=attn_mask, training=training)
124
+
125
+ # Unpack dense batch to original flat shape (static shapes: jit-friendly)
126
+ if batch is None:
127
+ h_global = attn_out[0]
128
+ else:
129
+ from k3_node.layers.aggr.base import from_dense_batch
130
+
131
+ h_global = from_dense_batch(attn_out, batch)
132
+
133
+ h_global = self.dropout(h_global, training=training)
134
+ h_global = h_global + x
135
+ if self.norm2 is not None:
136
+ h_global = self.norm2(h_global, training=training)
137
+ hs.append(h_global)
138
+
139
+ # Combine local and global
140
+ if len(hs) > 1:
141
+ out = hs[0] + hs[1]
142
+ else:
143
+ out = hs[0]
144
+
145
+ # MLP
146
+ mlp_h = self.dropout(ops.relu(self.mlp_l1(out)), training=training)
147
+ mlp_out = self.dropout(self.mlp_l2(mlp_h), training=training)
148
+ out = out + mlp_out
149
+
150
+ if self.norm3 is not None:
151
+ out = self.norm3(out, training=training)
152
+
153
+ return out
@@ -0,0 +1,262 @@
1
+ # ported from stellargraph
2
+ from keras import ops
3
+ from keras import activations, constraints, initializers, regularizers
4
+ from keras.layers import Layer, LeakyReLU, Dropout
5
+
6
+
7
+ class GraphAttention(Layer):
8
+ """
9
+ `k3_node.layers.GraphAttention`
10
+ Implementation of Graph Attention (GAT) layer
11
+
12
+ Args:
13
+ units: Positive integer, dimensionality of the output space.
14
+ attn_heads: Positive integer, number of attention heads.
15
+ attn_heads_reduction: {'concat', 'average'} Method for reducing attention heads.
16
+ in_dropout_rate: Dropout rate applied to the input (node features).
17
+ attn_dropout_rate: Dropout rate applied to attention coefficients.
18
+ activation: Activation function to use.
19
+ use_bias: Whether to add a bias to the linear transformation.
20
+ final_layer: Deprecated, use tf.gather or GatherIndices instead.
21
+ saliency_map_support: Whether to support saliency map calculations.
22
+ kernel_initializer: Initializer for the `kernel` weights matrix.
23
+ kernel_regularizer: Regularizer for the `kernel` weights matrix.
24
+ kernel_constraint: Constraint for the `kernel` weights matrix.
25
+ bias_initializer: Initializer for the bias vector.
26
+ bias_regularizer: Regularizer for the bias vector.
27
+ bias_constraint: Constraint for the bias vector.
28
+ attn_kernel_initializer: Initializer for the attention kernel weights matrix.
29
+ attn_kernel_regularizer: Regularizer for the attention kernel weights matrix.
30
+ attn_kernel_constraint: Constraint for the attention kernel weights matrix.
31
+ **kwargs: Additional arguments to pass to the `Layer` superclass.
32
+
33
+ Example:
34
+ ```python
35
+ import numpy as np
36
+ from k3_node.layers import GraphAttention
37
+
38
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
39
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
40
+
41
+ layer = GraphAttention(units=16, attn_heads=2)
42
+ out = layer(x, edge_index)
43
+ print(tuple(out.shape)) # (10, 2)
44
+ ```
45
+ """
46
+ def __init__(
47
+ self,
48
+ units,
49
+ attn_heads=1,
50
+ attn_heads_reduction="concat", # {'concat', 'average'}
51
+ in_dropout_rate=0.0,
52
+ attn_dropout_rate=0.0,
53
+ activation="relu",
54
+ use_bias=True,
55
+ final_layer=None,
56
+ saliency_map_support=False,
57
+ kernel_initializer="glorot_uniform",
58
+ kernel_regularizer=None,
59
+ kernel_constraint=None,
60
+ bias_initializer="zeros",
61
+ bias_regularizer=None,
62
+ bias_constraint=None,
63
+ attn_kernel_initializer="glorot_uniform",
64
+ attn_kernel_regularizer=None,
65
+ attn_kernel_constraint=None,
66
+ **kwargs,
67
+ ):
68
+ if attn_heads_reduction not in {"concat", "average"}:
69
+ raise ValueError(
70
+ "{}: Possible heads reduction methods: concat, average; received {}".format(
71
+ type(self).__name__, attn_heads_reduction
72
+ )
73
+ )
74
+
75
+ if isinstance(attn_heads, int) and attn_heads > 1 and "out_channels" not in kwargs:
76
+ self.in_channels = units
77
+ units = attn_heads
78
+ attn_heads = 1
79
+
80
+ self.units = units # Number of output features (F' in the paper)
81
+ self.attn_heads = attn_heads # Number of attention heads (K in the paper)
82
+ self.attn_heads_reduction = attn_heads_reduction # Eq. 5 and 6 in the paper
83
+ self.in_dropout_rate = in_dropout_rate # dropout rate for node features
84
+ self.attn_dropout_rate = attn_dropout_rate # dropout rate for attention coefs
85
+ self.activation = activations.get(activation) # Eq. 4 in the paper
86
+ self.use_bias = use_bias
87
+ if final_layer is not None:
88
+ raise ValueError(
89
+ "'final_layer' is not longer supported, use 'tf.gather' or 'GatherIndices' separately"
90
+ )
91
+
92
+ self.saliency_map_support = saliency_map_support
93
+ # Populated by build()
94
+ self.kernels = [] # Layer kernels for attention heads
95
+ self.biases = [] # Layer biases for attention heads
96
+ self.attn_kernels = [] # Attention kernels for attention heads
97
+
98
+ if attn_heads_reduction == "concat":
99
+ # Output will have shape (..., K * F')
100
+ self.output_dim = self.units * self.attn_heads
101
+ else:
102
+ # Output will have shape (..., F')
103
+ self.output_dim = self.units
104
+
105
+ self.kernel_initializer = initializers.get(kernel_initializer)
106
+ self.kernel_regularizer = regularizers.get(kernel_regularizer)
107
+ self.kernel_constraint = constraints.get(kernel_constraint)
108
+ self.bias_initializer = initializers.get(bias_initializer)
109
+ self.bias_regularizer = regularizers.get(bias_regularizer)
110
+ self.bias_constraint = constraints.get(bias_constraint)
111
+ self.attn_kernel_initializer = initializers.get(attn_kernel_initializer)
112
+ self.attn_kernel_regularizer = regularizers.get(attn_kernel_regularizer)
113
+ self.attn_kernel_constraint = constraints.get(attn_kernel_constraint)
114
+
115
+ super().__init__(**kwargs)
116
+ self.in_dropout = Dropout(in_dropout_rate) # created once, applied with `training`
117
+ self.attn_dropout = Dropout(attn_dropout_rate)
118
+
119
+ def build(self, input_shapes):
120
+ if isinstance(input_shapes, (list, tuple)) and len(input_shapes) > 0 and isinstance(input_shapes[0], (list, tuple)):
121
+ feat_shape = input_shapes[0]
122
+ else:
123
+ feat_shape = input_shapes
124
+ input_dim = int(feat_shape[-1]) if feat_shape is not None and feat_shape[-1] is not None else 8
125
+
126
+ # Variables to support integrated gradients
127
+ self.delta = self.add_weight(
128
+ name="ig_delta", shape=(), trainable=False, initializer=initializers.ones()
129
+ )
130
+ self.non_exist_edge = self.add_weight(
131
+ name="ig_non_exist_edge",
132
+ shape=(),
133
+ trainable=False,
134
+ initializer=initializers.zeros(),
135
+ )
136
+
137
+ # Initialize weights for each attention head
138
+ for head in range(self.attn_heads):
139
+ # Layer kernel
140
+ kernel = self.add_weight(
141
+ shape=(input_dim, self.units),
142
+ initializer=self.kernel_initializer,
143
+ regularizer=self.kernel_regularizer,
144
+ constraint=self.kernel_constraint,
145
+ name="kernel_{}".format(head),
146
+ )
147
+ self.kernels.append(kernel)
148
+
149
+ # # Layer bias
150
+ if self.use_bias:
151
+ bias = self.add_weight(
152
+ shape=(self.units,),
153
+ initializer=self.bias_initializer,
154
+ regularizer=self.bias_regularizer,
155
+ constraint=self.bias_constraint,
156
+ name="bias_{}".format(head),
157
+ )
158
+ self.biases.append(bias)
159
+
160
+ # Attention kernels
161
+ attn_kernel_self = self.add_weight(
162
+ shape=(self.units, 1),
163
+ initializer=self.attn_kernel_initializer,
164
+ regularizer=self.attn_kernel_regularizer,
165
+ constraint=self.attn_kernel_constraint,
166
+ name="attn_kernel_self_{}".format(head),
167
+ )
168
+ attn_kernel_neighs = self.add_weight(
169
+ shape=(self.units, 1),
170
+ initializer=self.attn_kernel_initializer,
171
+ regularizer=self.attn_kernel_regularizer,
172
+ constraint=self.attn_kernel_constraint,
173
+ name="attn_kernel_neigh_{}".format(head),
174
+ )
175
+ self.attn_kernels.append([attn_kernel_self, attn_kernel_neighs])
176
+ self.built = True
177
+
178
+ def call(self, inputs, A=None, training=None, **kwargs):
179
+ if A is not None:
180
+ X = inputs
181
+ elif isinstance(inputs, (list, tuple)):
182
+ X = inputs[0]
183
+ A = inputs[1]
184
+ else:
185
+ X, A = inputs, None
186
+
187
+ if A is not None and hasattr(A, "shape") and len(A.shape) == 2 and A.shape[0] == 2 and A.shape[1] != 2:
188
+ num_nodes = ops.shape(X)[-2]
189
+ a_dense = ops.zeros((num_nodes, num_nodes), dtype=X.dtype)
190
+ indices = ops.transpose(A, axes=[1, 0])
191
+ updates = ops.ones(shape=(ops.shape(A)[1],), dtype=X.dtype)
192
+ A = ops.scatter_update(a_dense, indices, updates)
193
+
194
+ assert len(ops.shape(A)) == 2, f"Adjacency matrix A should be 2-D"
195
+ N = ops.shape(A)[-1]
196
+
197
+ outputs = []
198
+ for head in range(self.attn_heads):
199
+ kernel = self.kernels[head] # W in the paper (F x F')
200
+ attention_kernel = self.attn_kernels[
201
+ head
202
+ ] # Attention kernel a in the paper (2F' x 1)
203
+
204
+ # Compute inputs to attention network
205
+
206
+ features = ops.dot(X, kernel) # (N x F')
207
+
208
+ # Compute feature combinations
209
+ # Note: [[a_1], [a_2]]^T [[Wh_i], [Wh_2]] = [a_1]^T [Wh_i] + [a_2]^T [Wh_j]
210
+ attn_for_self = ops.dot(
211
+ features, attention_kernel[0]
212
+ ) # (N x 1), [a_1]^T [Wh_i]
213
+ attn_for_neighs = ops.dot(
214
+ features, attention_kernel[1]
215
+ ) # (N x 1), [a_2]^T [Wh_j]
216
+
217
+ # Attention head a(Wh_i, Wh_j) = a^T [[Wh_i], [Wh_j]]
218
+ dense = attn_for_self + ops.transpose(
219
+ attn_for_neighs
220
+ ) # (N x N) via broadcasting
221
+
222
+ dense = LeakyReLU(0.2)(dense)
223
+
224
+ if not self.saliency_map_support:
225
+ mask = -10e9 * (1.0 - A)
226
+ dense += mask
227
+ dense = ops.softmax(dense) # (N x N), Eq. 3 of the paper
228
+
229
+ else:
230
+ # dense = dense - tf.reduce_max(dense)
231
+ # GAT with support for saliency calculations
232
+ W = (self.delta * A) * ops.exp(
233
+ dense - ops.max(dense, axis=1, keepdims=True)
234
+ ) * (1 - self.non_exist_edge) + self.non_exist_edge * (
235
+ A + self.delta * (ops.ones((N, N)) - A) + ops.eye(N)
236
+ ) * ops.exp(
237
+ dense - ops.max(dense, axis=1, keepdims=True)
238
+ )
239
+ dense = W / ops.sum(W, axis=1, keepdims=True)
240
+
241
+ # Apply dropout to features and attention coefficients
242
+ dropout_feat = self.in_dropout(features, training=training) # (N x F')
243
+ dropout_attn = self.attn_dropout(dense, training=training) # (N x N)
244
+
245
+ # Linear combination with neighbors' features [YT: see Eq. 4]
246
+ node_features = ops.dot(dropout_attn, dropout_feat) # (N x F')
247
+
248
+ if self.use_bias:
249
+ node_features = ops.add(node_features, self.biases[head])
250
+
251
+ # Add output of attention head to final output
252
+ outputs.append(node_features)
253
+
254
+ # Aggregate the heads' output according to the reduction method
255
+ if self.attn_heads_reduction == "concat":
256
+ output = ops.concatenate(outputs, axis=1) # (N x KF')
257
+ else:
258
+ output = ops.mean(ops.stack(outputs), axis=0) # N x F')
259
+
260
+ output = self.activation(output)
261
+
262
+ return output
@@ -0,0 +1,84 @@
1
+ from typing import Union, Tuple
2
+ from keras import layers, ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+
6
+
7
+ class GraphConv(MessagePassing):
8
+ r"""The graph neural network operator from the `"Weisfeiler and Leman Go
9
+ Neural: Higher-order Graph Neural Networks"
10
+ <https://arxiv.org/abs/1810.02244>`_ 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
+ aggr: The aggregation scheme to use (``"add"``, ``"mean"``, ``"max"``).
16
+ (default: ``"add"``)
17
+ bias: If set to :obj:`False`, the layer will not learn
18
+ an additive bias. (default: ``"True"``)
19
+
20
+ Example:
21
+ ```python
22
+ import numpy as np
23
+ from k3_node.layers import GraphConv
24
+
25
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
26
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
27
+
28
+ layer = GraphConv(in_channels=8, out_channels=16)
29
+ out = layer(x, edge_index)
30
+ print(tuple(out.shape)) # (10, 16)
31
+ ```
32
+ """
33
+
34
+ weighted_sum_message = True
35
+
36
+ def __init__(
37
+ self,
38
+ in_channels: Union[int, Tuple[int, int]],
39
+ out_channels: int,
40
+ aggr: str = "add",
41
+ bias: bool = True,
42
+ **kwargs,
43
+ ):
44
+ super().__init__(aggr=aggr, **kwargs)
45
+ self.in_channels = in_channels
46
+ self.out_channels = out_channels
47
+ self.use_bias = bias
48
+
49
+ self.lin_rel = layers.Dense(out_channels, use_bias=bias)
50
+ self.lin_root = layers.Dense(out_channels, use_bias=False)
51
+
52
+ def build(self, input_shape):
53
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
54
+ in_channels_l = input_shape[0][-1]
55
+ in_channels_r = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_l
56
+ else:
57
+ in_channels_l = input_shape[-1]
58
+ in_channels_r = input_shape[-1]
59
+
60
+ self.lin_rel.build((None, in_channels_l))
61
+ self.lin_root.build((None, in_channels_r))
62
+ self.built = True
63
+
64
+ def call(self, x, edge_index=None, edge_weight=None, size=None, **kwargs):
65
+ if edge_index is None and isinstance(x, (tuple, list)):
66
+ x, edge_index = x[0], x[1]
67
+
68
+ if not isinstance(x, (tuple, list)):
69
+ x_src = x
70
+ x_dst = x
71
+ else:
72
+ x_src, x_dst = x[0], x[1]
73
+
74
+ out = self.propagate(edge_index, x=(x_src, x_dst), edge_weight=edge_weight, size=size)
75
+ out = self.lin_rel(out)
76
+ if x_dst is not None:
77
+ out = out + self.lin_root(x_dst)
78
+ return out
79
+
80
+ def message(self, x_j, edge_weight=None):
81
+ if edge_weight is None:
82
+ return x_j
83
+ return ops.expand_dims(edge_weight, -1) * x_j
84
+
@@ -0,0 +1,93 @@
1
+ from typing import Optional, Union, Tuple
2
+ import keras
3
+ from keras import ops
4
+ from keras.layers import Dense
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+ from k3_node.layers.pool.knn import knn_graph
8
+
9
+
10
+ class GravNetConv(MessagePassing):
11
+ r"""The GravNet operator from the `"Learning Representations of Irregular
12
+ Particle-Detector Geometry with Distance-Weighted Graph Networks"
13
+ <https://arxiv.org/abs/1902.07987>`_ paper.
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import GravNetConv
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
+ # Neighbors are found by k-NN in a learned space, so no edge_index is needed
24
+ layer = GravNetConv(in_channels=8, out_channels=16, space_dimensions=3, propagate_dimensions=4, k=3)
25
+ out = layer(x)
26
+ print(tuple(out.shape)) # (10, 16)
27
+ ```
28
+ """
29
+ def __init__(
30
+ self,
31
+ in_channels: int,
32
+ out_channels: int,
33
+ space_dimensions: int,
34
+ propagate_dimensions: int,
35
+ k: int,
36
+ num_workers: Optional[int] = None,
37
+ **kwargs,
38
+ ):
39
+ kwargs.setdefault("aggr", "mean")
40
+ super().__init__(node_dim=0, **kwargs)
41
+
42
+ self.in_channels = in_channels
43
+ self.out_channels = out_channels
44
+ self.space_dimensions = space_dimensions
45
+ self.propagate_dimensions = propagate_dimensions
46
+ self.k = k
47
+
48
+ self.lin_s = Dense(space_dimensions, use_bias=True)
49
+ self.lin_h = Dense(propagate_dimensions, use_bias=True)
50
+ self.lin_out1 = Dense(out_channels, use_bias=True)
51
+ self.lin_out2 = Dense(out_channels, use_bias=True)
52
+
53
+ def build(self, input_shape=None):
54
+ self.lin_s.build((None, self.in_channels))
55
+ self.lin_h.build((None, self.in_channels))
56
+ self.lin_out1.build((None, self.in_channels))
57
+ self.lin_out2.build((None, self.propagate_dimensions))
58
+ self.built = True
59
+
60
+ def call(self, inputs, edge_index=None, **kwargs):
61
+ if isinstance(inputs, (list, tuple)):
62
+ x = inputs[0]
63
+ else:
64
+ x = inputs
65
+
66
+ if not self.built:
67
+ self.build()
68
+
69
+ s = self.lin_s(x)
70
+ h = self.lin_h(x)
71
+
72
+ if edge_index is None:
73
+ edge_index = knn_graph(s, k=self.k)
74
+
75
+ edge_index = ops.cast(edge_index, "int32")
76
+ s_src = ops.take(s, edge_index[0], axis=0)
77
+ s_dst = ops.take(s, edge_index[1], axis=0)
78
+ dist_sq = ops.sum(ops.square(s_src - s_dst), axis=-1)
79
+ edge_weight = ops.exp(-10.0 * dist_sq)
80
+
81
+ num_nodes = ops.shape(x)[0]
82
+ out = self.propagate(
83
+ edge_index,
84
+ x=h,
85
+ edge_weight=edge_weight,
86
+ size=(num_nodes, num_nodes),
87
+ )
88
+
89
+ return self.lin_out1(x) + self.lin_out2(out)
90
+
91
+ def message(self, x_j, edge_weight):
92
+ return ops.expand_dims(edge_weight, -1) * x_j
93
+