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,451 @@
1
+ import keras
2
+ import inspect
3
+ from typing import Any, List, Optional, Tuple, Union
4
+ from keras import layers, ops
5
+
6
+ from k3_node.utils import (
7
+ is_layer_kwarg,
8
+ deserialize_kwarg,
9
+ serialize_kwarg,
10
+ is_keras_kwarg,
11
+ deserialize_scatter,
12
+ serialize_scatter,
13
+ )
14
+ from k3_node.ops import get_source_target
15
+ from k3_node.layers.aggr.resolver import aggregation_resolver
16
+ from k3_node.layers.aggr.base import Aggregation
17
+ from k3_node.layers.aggr.basic import SumAggregation
18
+ from k3_node.ops.segment import segment_sum
19
+
20
+
21
+ class MessagePassing(layers.Layer):
22
+ r"""Base class for creating Message Passing Neural Networks (MPNNs).
23
+
24
+ Args:
25
+ aggr: The aggregation scheme to use, such as ``"add"``, ``"sum"``,
26
+ ``"mean"``, ``"min"``, ``"max"``, ``"mul"``, or an instance of
27
+ :class:`~k3_node.layers.aggr.Aggregation`. (default: ``"add"``)
28
+ flow: The direction of message passing (``"source_to_target"`` or
29
+ ``"target_to_source"``). (default: ``"source_to_target"``)
30
+ node_dim: The axis along which to index node features. (default: ``-2``)
31
+ decomposed_layers: Number of decomposed layers for memory-efficient
32
+ aggregation. (default: ``1``)
33
+
34
+ Example:
35
+ ```python
36
+ import numpy as np
37
+ import keras
38
+ from k3_node.layers import MessagePassing
39
+
40
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
41
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
42
+
43
+ # A minimal custom layer: average the neighbors' features, then transform them.
44
+ class MeanConv(MessagePassing):
45
+ def __init__(self, units, **kwargs):
46
+ super().__init__(aggr="mean", **kwargs)
47
+ self.dense = keras.layers.Dense(units)
48
+
49
+ def call(self, x, edge_index):
50
+ return self.dense(self.propagate(edge_index, x=x))
51
+
52
+ def message(self, x_j):
53
+ return x_j # features of the source node of every edge
54
+
55
+
56
+ layer = MeanConv(16)
57
+ out = layer(x, edge_index)
58
+ print(tuple(out.shape)) # (10, 16)
59
+ ```
60
+ """
61
+
62
+ def __init__(
63
+ self,
64
+ aggr: Union[str, List[str], Aggregation, None] = "add",
65
+ flow: str = "source_to_target",
66
+ node_dim: int = -2,
67
+ decomposed_layers: int = 1,
68
+ **kwargs,
69
+ ):
70
+ # Support legacy aggregate arg
71
+ if "aggregate" in kwargs:
72
+ aggr = kwargs.pop("aggregate")
73
+
74
+ # Extract and set layer kwargs for backwards compatibility
75
+ self.kwargs_keys = []
76
+ for key in list(kwargs.keys()):
77
+ if is_layer_kwarg(key):
78
+ attr = kwargs.pop(key)
79
+ attr = deserialize_kwarg(key, attr)
80
+ self.kwargs_keys.append(key)
81
+ setattr(self, key, attr)
82
+
83
+ unknown = sorted(k for k in kwargs if not is_keras_kwarg(k))
84
+ if unknown:
85
+ # Silently dropping these hides typos such as `num_heads=` for `heads=`.
86
+ raise TypeError(f"{type(self).__name__}() got unexpected keyword argument(s): {', '.join(unknown)}")
87
+ super().__init__(**kwargs)
88
+ self.aggr = aggr
89
+ self.flow = flow
90
+ self.node_dim = node_dim
91
+ self.decomposed_layers = decomposed_layers
92
+
93
+ if flow not in ["source_to_target", "target_to_source"]:
94
+ raise ValueError(f"Flow {flow} must be 'source_to_target' or 'target_to_source'")
95
+
96
+ # Legacy scatter operator support
97
+ try:
98
+ scatter_name = "sum" if aggr in ("add", "sum") else aggr
99
+ self.agg = deserialize_scatter(scatter_name)
100
+ except Exception:
101
+ self.agg = ops.segment_sum
102
+
103
+ if aggr is None:
104
+ self.aggr_module = None
105
+ elif isinstance(aggr, Aggregation):
106
+ self.aggr_module = aggr
107
+ elif isinstance(aggr, str):
108
+ aggr_name = "sum" if aggr == "add" else aggr
109
+ try:
110
+ self.aggr_module = aggregation_resolver(aggr_name)
111
+ except Exception:
112
+ self.aggr_module = SumAggregation()
113
+ elif isinstance(aggr, (list, tuple)):
114
+ from k3_node.layers.aggr.multi import MultiAggregation
115
+ self.aggr_module = MultiAggregation(aggr)
116
+ else:
117
+ self.aggr_module = SumAggregation()
118
+
119
+ self.msg_signature = inspect.signature(self.message).parameters
120
+ self.agg_signature = inspect.signature(self.aggregate).parameters
121
+ self.upd_signature = inspect.signature(self.update).parameters
122
+
123
+ def build(self, input_shape=None):
124
+ self.built = True
125
+
126
+ @staticmethod
127
+ def get_inputs(inputs):
128
+ if len(inputs) == 3:
129
+ x, a, e = inputs
130
+ if hasattr(e, "shape") and e.shape is not None:
131
+ assert len(e.shape) in (2, 3), "E must have rank 2 or 3"
132
+ elif len(inputs) == 2:
133
+ x, a = inputs
134
+ e = None
135
+ else:
136
+ raise ValueError(
137
+ "Expected 2 or 3 inputs tensors (X, A, E), got {}.".format(len(inputs))
138
+ )
139
+ if hasattr(a, "shape") and a.shape is not None:
140
+ assert len(a.shape) == 2, "A must have rank 2"
141
+ return x, a, e
142
+
143
+ def get_targets(self, x):
144
+ return ops.take(x, self.index_targets, axis=self.node_dim)
145
+
146
+ def get_sources(self, x):
147
+ return ops.take(x, self.index_sources, axis=self.node_dim)
148
+
149
+ def get_kwargs(self, x, a, e, signature, kwargs):
150
+ output = {}
151
+ for k in signature.keys():
152
+ if k == "kwargs":
153
+ pass
154
+ elif k == "x":
155
+ output[k] = x
156
+ elif k == "a":
157
+ output[k] = a
158
+ elif k == "e":
159
+ output[k] = e
160
+ elif k in kwargs:
161
+ output[k] = kwargs[k]
162
+ elif signature[k].default is inspect.Parameter.empty:
163
+ pass
164
+ else:
165
+ pass
166
+ return output
167
+
168
+ def _get_dim_size(self, kwargs, i, size=None):
169
+ if size is not None and size[1] is not None:
170
+ return size[1]
171
+ x = kwargs.get("x", None)
172
+ if x is not None:
173
+ if isinstance(x, (tuple, list)):
174
+ target_x = x[1] if x[1] is not None else x[0]
175
+ if target_x is not None:
176
+ if hasattr(target_x, "shape") and target_x.shape[self.node_dim] is not None:
177
+ return int(target_x.shape[self.node_dim])
178
+ return ops.shape(target_x)[self.node_dim]
179
+ else:
180
+ if hasattr(x, "shape") and x.shape[self.node_dim] is not None:
181
+ return int(x.shape[self.node_dim])
182
+ return ops.shape(x)[self.node_dim]
183
+ for k, val in kwargs.items():
184
+ if k in ("edge_index", "edge_attr", "edge_weight", "ptr"):
185
+ continue
186
+ if hasattr(val, "shape") and len(val.shape) >= 2:
187
+ if val.shape[self.node_dim] is not None:
188
+ return int(val.shape[self.node_dim])
189
+ return ops.shape(val)[self.node_dim]
190
+ if i is not None:
191
+ from k3_node.layers.conv.utils import is_tracing
192
+ if is_tracing(i):
193
+ return 0
194
+ if hasattr(i, "shape") and len(i.shape) > 0 and i.shape[0] == 0:
195
+ return 0
196
+ try:
197
+ return int(ops.max(i)) + 1
198
+ except Exception:
199
+ return 0
200
+ return 0
201
+
202
+ def propagate(self, *args, **kwargs: Any):
203
+ r"""The initial call to start propagating messages."""
204
+ # Detect legacy Spektral-style call: propagate(x, a, e=None, **kwargs)
205
+ is_spektral = False
206
+ if len(args) >= 2:
207
+ arg0, arg1 = args[0], args[1]
208
+ shape0 = arg0.shape if hasattr(arg0, "shape") and arg0.shape is not None else ()
209
+ shape1 = arg1.shape if hasattr(arg1, "shape") and arg1.shape is not None else ()
210
+ if len(shape1) == 2 and shape1[0] is not None and shape1[0] == shape1[1]:
211
+ is_spektral = True
212
+ elif len(shape0) >= 2 and shape0[0] is not None and shape0[0] != 2 and not isinstance(arg1, (tuple, list, type(None))):
213
+ if len(shape1) == 2 and shape1[0] is not None and shape1[0] != 2:
214
+ is_spektral = True
215
+
216
+ if is_spektral:
217
+ x, a = args[0], args[1]
218
+ e = args[2] if len(args) >= 3 else kwargs.get("e", None)
219
+ self.n_nodes = x.shape[-2] if hasattr(x, "shape") and x.shape[-2] is not None else ops.shape(x)[-2]
220
+ self.index_sources, self.index_targets = get_source_target(a)
221
+
222
+ # Call legacy message
223
+ msg_kwargs = self.get_kwargs(x, a, e, self.msg_signature, kwargs)
224
+ messages = self.message(**msg_kwargs)
225
+
226
+ # Call legacy aggregate
227
+ agg_kwargs = self.get_kwargs(x, a, e, self.agg_signature, kwargs)
228
+ embeddings = self.aggregate(messages, **agg_kwargs)
229
+
230
+ # Call legacy update
231
+ upd_kwargs = self.get_kwargs(x, a, e, self.upd_signature, kwargs)
232
+ output = self.update(embeddings, **upd_kwargs)
233
+ return output
234
+
235
+ # Standard PyG MessagePassing propagate(edge_index, size=None, **kwargs)
236
+ if len(args) >= 1:
237
+ edge_index = args[0]
238
+ size = args[1] if len(args) >= 2 else kwargs.pop("size", None)
239
+ else:
240
+ edge_index = kwargs.pop("edge_index")
241
+ size = kwargs.pop("size", None)
242
+
243
+ edge_index = ops.convert_to_tensor(edge_index)
244
+
245
+ # Handle dense adjacency [N, N] passed as edge_index
246
+ e_shape = getattr(edge_index, "shape", None)
247
+ if e_shape is not None and len(e_shape) == 2 and e_shape[0] is not None and e_shape[1] is not None and e_shape[0] > 2 and e_shape[0] == e_shape[1]:
248
+ where_adj = ops.where(edge_index != 0)
249
+ where_adj = where_adj if not isinstance(where_adj, list) else where_adj
250
+ edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
251
+
252
+ if self.flow == "source_to_target":
253
+ i = edge_index[1] # Target
254
+ j = edge_index[0] # Source
255
+ else:
256
+ i = edge_index[0]
257
+ j = edge_index[1]
258
+
259
+ i = ops.cast(i, "int32")
260
+ j = ops.cast(j, "int32")
261
+ self.index_targets = i
262
+ self.index_sources = j
263
+
264
+ dim_size = size[1] if size is not None and size[1] is not None else self._get_dim_size(kwargs, i, size)
265
+ self.n_nodes = dim_size
266
+
267
+ # Layers whose message is `edge_weight * x_j`, summed, can skip per-edge messages
268
+ fused = self._fused_weighted_sum(i, j, dim_size, size, kwargs)
269
+ if fused is not None:
270
+ return self.update(fused, **self._update_kwargs(kwargs))
271
+
272
+ # Construct message arguments
273
+ msg_kwargs = {}
274
+ for param_name in self.msg_signature.keys():
275
+ if param_name in kwargs:
276
+ msg_kwargs[param_name] = kwargs[param_name]
277
+ elif param_name in ("dim_size", "size_i"):
278
+ msg_kwargs[param_name] = dim_size
279
+ elif param_name == "size_j":
280
+ msg_kwargs["size_j"] = size[0] if size is not None and size[0] is not None else dim_size
281
+ elif param_name == "index":
282
+ msg_kwargs["index"] = i
283
+ elif param_name == "edge_index":
284
+ msg_kwargs["edge_index"] = edge_index
285
+ elif param_name.endswith("_i"):
286
+ root = param_name[:-2]
287
+ if root in kwargs:
288
+ val = kwargs[root]
289
+ val = val[1] if isinstance(val, (tuple, list)) else val
290
+ msg_kwargs[param_name] = ops.take(val, i, axis=self.node_dim) if val is not None else None
291
+ elif param_name.endswith("_j"):
292
+ root = param_name[:-2]
293
+ if root in kwargs:
294
+ val = kwargs[root]
295
+ val = val[0] if isinstance(val, (tuple, list)) else val
296
+ msg_kwargs[param_name] = ops.take(val, j, axis=self.node_dim) if val is not None else None
297
+ elif param_name == "index":
298
+ msg_kwargs["index"] = i
299
+ elif param_name in ("dim_size", "size_i"):
300
+ msg_kwargs[param_name] = dim_size
301
+ elif param_name == "size_j":
302
+ msg_kwargs["size_j"] = size[0] if size is not None and size[0] is not None else dim_size
303
+ elif param_name == "edge_index":
304
+ msg_kwargs["edge_index"] = edge_index
305
+
306
+ out = self.message(**msg_kwargs)
307
+
308
+ # Aggregate
309
+ agg_kwargs = {}
310
+ for param_name in self.agg_signature.keys():
311
+ if param_name in kwargs and param_name not in ["inputs", "index", "ptr", "dim_size"]:
312
+ agg_kwargs[param_name] = kwargs[param_name]
313
+
314
+ out = self.aggregate(
315
+ out,
316
+ index=i,
317
+ ptr=kwargs.get("ptr", None),
318
+ dim_size=dim_size,
319
+ **agg_kwargs,
320
+ )
321
+
322
+ return self.update(out, **self._update_kwargs(kwargs))
323
+
324
+ def _update_kwargs(self, kwargs):
325
+ upd_kwargs = {}
326
+ for param_name in self.upd_signature.keys():
327
+ if param_name in kwargs and param_name not in ["inputs", "aggr_out", "embeddings"]:
328
+ upd_kwargs[param_name] = kwargs[param_name]
329
+ elif param_name.endswith("_i"):
330
+ root = param_name[:-2]
331
+ if root in kwargs:
332
+ val = kwargs[root]
333
+ val = val[1] if isinstance(val, (tuple, list)) else val
334
+ upd_kwargs[param_name] = val
335
+ return upd_kwargs
336
+
337
+ #: Set to ``True`` in layers whose ``message`` is ``edge_weight * x_j`` (or ``x_j``) and that
338
+ #: sum the messages; ``propagate`` then uses a sparse matrix product (see ``k3_node.ops.spmm``).
339
+ weighted_sum_message = False
340
+
341
+ def _fused_weighted_sum(self, i, j, dim_size, size, kwargs):
342
+ if not self.weighted_sum_message or keras.config.backend() != "torch":
343
+ return None
344
+ if not (self.aggr in ("add", "sum") and type(self.aggr_module).__name__ == "SumAggregation"):
345
+ return None
346
+ x = kwargs.get("x")
347
+ if x is None or isinstance(x, (tuple, list)) or len(getattr(x, "shape", ())) != 2 or self.node_dim not in (0, -2):
348
+ return None
349
+ if set(kwargs) - {"x", "edge_weight"}:
350
+ return None
351
+ from k3_node.ops.sparse import spmm
352
+
353
+ return spmm(i, j, kwargs.get("edge_weight"), x, dim_size)
354
+
355
+ def message(self, x=None, x_j=None, **kwargs):
356
+ r"""Constructs messages from node :math:`j` to node :math:`i`."""
357
+ if x_j is not None:
358
+ return x_j
359
+ if x is not None and hasattr(self, "index_sources"):
360
+ return self.get_sources(x)
361
+ return x_j
362
+
363
+ def aggregate(
364
+ self,
365
+ inputs=None,
366
+ index=None,
367
+ ptr: Optional[Any] = None,
368
+ dim_size: Optional[int] = None,
369
+ **kwargs,
370
+ ):
371
+ r"""Aggregates messages from neighbors as given by :obj:`index`."""
372
+ # Check if called as Spektral: aggregate(messages, **kwargs)
373
+ if inputs is not None and index is None and hasattr(self, "index_targets"):
374
+ index = self.index_targets
375
+ dim_size = self.n_nodes
376
+ if hasattr(self, "agg") and callable(self.agg):
377
+ return self.agg(inputs, index, dim_size)
378
+
379
+ if self.aggr_module is not None:
380
+ return self.aggr_module(
381
+ inputs, index=index, ptr=ptr, dim_size=dim_size, dim=self.node_dim
382
+ )
383
+ return segment_sum(inputs, index, num_segments=dim_size)
384
+
385
+ def update(self, embeddings=None, **kwargs):
386
+ r"""Updates node embeddings."""
387
+ return embeddings
388
+
389
+ def edge_updater(
390
+ self,
391
+ edge_index,
392
+ size: Optional[Tuple[Optional[int], Optional[int]]] = None,
393
+ **kwargs: Any,
394
+ ):
395
+ r"""Computes or updates edge-level representations."""
396
+ edge_index = ops.convert_to_tensor(edge_index)
397
+ if self.flow == "source_to_target":
398
+ i = edge_index[1]
399
+ j = edge_index[0]
400
+ else:
401
+ i = edge_index[0]
402
+ j = edge_index[1]
403
+
404
+ i = ops.cast(i, "int32")
405
+ j = ops.cast(j, "int32")
406
+
407
+ dim_size = self._get_dim_size(kwargs, i, size)
408
+ edge_update_params = inspect.signature(self.edge_update).parameters
409
+
410
+ upd_kwargs = {}
411
+ for param_name in edge_update_params.keys():
412
+ if param_name in kwargs:
413
+ upd_kwargs[param_name] = kwargs[param_name]
414
+ elif param_name.endswith("_i"):
415
+ root = param_name[:-2]
416
+ if root in kwargs:
417
+ val = kwargs[root]
418
+ val = val[1] if isinstance(val, (tuple, list)) else val
419
+ upd_kwargs[param_name] = ops.take(val, i, axis=self.node_dim) if val is not None else None
420
+ elif param_name.endswith("_j"):
421
+ root = param_name[:-2]
422
+ if root in kwargs:
423
+ val = kwargs[root]
424
+ val = val[0] if isinstance(val, (tuple, list)) else val
425
+ upd_kwargs[param_name] = ops.take(val, j, axis=self.node_dim) if val is not None else None
426
+ elif param_name == "index":
427
+ upd_kwargs["index"] = i
428
+ elif param_name == "dim_size":
429
+ upd_kwargs["dim_size"] = dim_size
430
+ elif param_name == "edge_index":
431
+ upd_kwargs["edge_index"] = edge_index
432
+
433
+ return self.edge_update(**upd_kwargs)
434
+
435
+ def edge_update(self, **kwargs):
436
+ r"""Computes or updates edge attributes."""
437
+ raise NotImplementedError
438
+
439
+ def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
440
+ r"""Default call handler supporting both (x, edge_index) and legacy (inputs,) tuples."""
441
+ if edge_index is None and isinstance(inputs, (tuple, list)):
442
+ x, a, e = self.get_inputs(inputs)
443
+ return self.propagate(x, a, e, **kwargs)
444
+ if edge_index is not None:
445
+ if edge_attr is not None:
446
+ kwargs["edge_attr"] = edge_attr
447
+ kwargs["e"] = edge_attr
448
+ return self.propagate(edge_index, x=inputs, **kwargs)
449
+ raise NotImplementedError(
450
+ f"Layer {self.__class__.__name__} does not implement call() with inputs={inputs}, edge_index={edge_index}"
451
+ )
@@ -0,0 +1,95 @@
1
+ from typing import Union, Tuple
2
+ import keras
3
+ from keras import layers, ops
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.conv.utils import degree
7
+
8
+
9
+ class MFConv(MessagePassing):
10
+ r"""The molecular fingerprint graph convolutional operator from the
11
+ `"Convolutional Networks on Graphs for Learning Molecular Fingerprints"
12
+ <https://arxiv.org/abs/1509.09292>`_ paper.
13
+
14
+ Args:
15
+ in_channels: Size of each input sample, or a tuple for bipartite graphs.
16
+ out_channels: Size of each output sample.
17
+ max_degree: The maximum degree of any node. (default: ``10``)
18
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
19
+ (default: ``True``)
20
+
21
+ Example:
22
+ ```python
23
+ import numpy as np
24
+ from k3_node.layers import MFConv
25
+
26
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
27
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
28
+
29
+ layer = MFConv(in_channels=8, out_channels=16)
30
+ out = layer(x, edge_index)
31
+ print(tuple(out.shape)) # (10, 16)
32
+ ```
33
+ """
34
+
35
+ def __init__(
36
+ self,
37
+ in_channels: Union[int, Tuple[int, int]],
38
+ out_channels: int,
39
+ max_degree: int = 10,
40
+ bias: bool = True,
41
+ **kwargs,
42
+ ):
43
+ super().__init__(aggr="add", **kwargs)
44
+ self.in_channels = in_channels
45
+ self.out_channels = out_channels
46
+ self.max_degree = max_degree
47
+ self.use_bias = bias
48
+
49
+ self.lins_l = [layers.Dense(out_channels, use_bias=bias) for _ in range(max_degree + 1)]
50
+ self.lins_r = [layers.Dense(out_channels, use_bias=False) for _ in range(max_degree + 1)]
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_src = input_shape[0][-1]
55
+ in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
56
+ else:
57
+ in_channels_src = input_shape[-1]
58
+ in_channels_dst = input_shape[-1]
59
+
60
+ for lin_l in self.lins_l:
61
+ lin_l.build((None, in_channels_src))
62
+ for lin_r in self.lins_r:
63
+ lin_r.build((None, in_channels_dst))
64
+ self.built = True
65
+
66
+ def call(self, x, edge_index=None, size=None, **kwargs):
67
+ if edge_index is None and isinstance(x, (tuple, list)):
68
+ x, edge_index = x[0], x[1]
69
+
70
+ if not isinstance(x, (tuple, list)):
71
+ x_src, x_dst = x, x
72
+ else:
73
+ x_src, x_dst = x[0], x[1]
74
+
75
+ target_idx = edge_index[1] if self.flow == "source_to_target" else edge_index[0]
76
+ N = ops.shape(x_dst)[self.node_dim] if x_dst is not None else ops.shape(x_src)[self.node_dim]
77
+ deg = degree(target_idx, num_nodes=N)
78
+ deg = ops.clip(deg, 0, self.max_degree)
79
+
80
+ h = self.propagate(edge_index, x=(x_src, x_dst), size=size)
81
+
82
+ out = ops.zeros((ops.shape(h)[0], self.out_channels), dtype=h.dtype)
83
+ for i, (lin_l, lin_r) in enumerate(zip(self.lins_l, self.lins_r)):
84
+ mask = ops.equal(deg, i)
85
+ mask = ops.expand_dims(ops.cast(mask, h.dtype), -1)
86
+ term = lin_l(h)
87
+ if x_dst is not None:
88
+ term = term + lin_r(x_dst)
89
+ out = out + mask * term
90
+
91
+ return out
92
+
93
+ def message(self, x_j):
94
+ return x_j
95
+
@@ -0,0 +1,108 @@
1
+ from typing import Optional, List
2
+
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.conv.utils import gcn_norm
8
+
9
+
10
+ class MixHopConv(MessagePassing):
11
+ r"""The MixHop graph convolutional operator from the
12
+ `"Higher-Order Graph Convolutional Networks via MixHop"
13
+ <https://arxiv.org/abs/1905.00067>`_ paper.
14
+
15
+ Example:
16
+ ```python
17
+ import numpy as np
18
+ from k3_node.layers import MixHopConv
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
+ layer = MixHopConv(in_channels=8, out_channels=16, powers=[0, 1, 2])
24
+ out = layer(x, edge_index)
25
+ print(tuple(out.shape)) # (10, 48)
26
+ ```
27
+ """
28
+
29
+ weighted_sum_message = True
30
+ def __init__(
31
+ self,
32
+ in_channels: int,
33
+ out_channels: int,
34
+ powers: Optional[List[int]] = None,
35
+ add_self_loops: bool = True,
36
+ bias: bool = True,
37
+ **kwargs,
38
+ ):
39
+ kwargs.setdefault("aggr", "add")
40
+ super().__init__(**kwargs)
41
+
42
+ if powers is None:
43
+ powers = [0, 1, 2]
44
+
45
+ self.in_channels = in_channels
46
+ self.out_channels = out_channels
47
+ self.powers = powers
48
+ self.add_self_loops = add_self_loops
49
+ self.use_bias = bias
50
+
51
+ self.lins = [Dense(out_channels, use_bias=False) for _ in range(max(powers) + 1)]
52
+
53
+ if bias:
54
+ self.bias = self.add_weight(
55
+ shape=(len(powers) * out_channels,),
56
+ initializer="zeros",
57
+ name="bias",
58
+ )
59
+ else:
60
+ self.bias = None
61
+
62
+ def build(self, input_shape=None):
63
+ for lin in self.lins:
64
+ lin.build((None, self.in_channels))
65
+ self.built = True
66
+
67
+ def call(self, inputs, edge_index=None, edge_weight=None, **kwargs):
68
+ if edge_index is None:
69
+ if isinstance(inputs, (list, tuple)):
70
+ if len(inputs) == 3:
71
+ x, edge_index, edge_weight = inputs
72
+ elif len(inputs) == 2:
73
+ x, edge_index = inputs
74
+ else:
75
+ raise ValueError(f"Unexpected input length {len(inputs)}")
76
+ else:
77
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
78
+ else:
79
+ x = inputs
80
+
81
+ if not self.built:
82
+ self.build()
83
+
84
+ num_nodes = ops.shape(x)[0]
85
+ edge_index, edge_weight = gcn_norm(
86
+ edge_index,
87
+ edge_weight,
88
+ num_nodes=num_nodes,
89
+ add_self_loops=self.add_self_loops,
90
+ dtype=x.dtype,
91
+ )
92
+
93
+ outs = [self.lins[0](x)]
94
+ curr_x = x
95
+ for lin in self.lins[1:]:
96
+ curr_x = self.propagate(edge_index, x=curr_x, edge_weight=edge_weight)
97
+ outs.append(lin(curr_x))
98
+
99
+ out = ops.concatenate([outs[p] for p in self.powers], axis=-1)
100
+
101
+ if self.bias is not None:
102
+ out = out + self.bias
103
+
104
+ return out
105
+
106
+ def message(self, x_j, edge_weight=None):
107
+ return x_j if edge_weight is None else ops.expand_dims(edge_weight, -1) * x_j
108
+