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,140 @@
1
+ # ported from spektral
2
+
3
+ from keras import ops
4
+ from keras.layers import Dense
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+
8
+
9
+ class CrystalConv(MessagePassing):
10
+ """
11
+ `k3_node.layers.CrystalConv`
12
+ Implementation of Crystal Graph Convolutional Neural Networks (CGCNN) layer
13
+
14
+ Args:
15
+ channels: The number of channels/units (optional).
16
+ edge_dim: The dimensionality of edge features (optional).
17
+ aggregate: Aggregation function to use (one of 'sum', 'mean', 'max').
18
+ activation: Activation function to use.
19
+ use_bias: Whether to add a bias to the linear transformation.
20
+ kernel_initializer: Initializer for the `kernel` weights matrix.
21
+ bias_initializer: Initializer for the bias vector.
22
+ kernel_regularizer: Regularizer for the `kernel` weights matrix.
23
+ bias_regularizer: Regularizer for the bias vector.
24
+ activity_regularizer: Regularizer for the output.
25
+ kernel_constraint: Constraint for the `kernel` weights matrix.
26
+ bias_constraint: Constraint for the bias vector.
27
+ **kwargs: Additional arguments to pass to the `MessagePassing` superclass.
28
+
29
+ Example:
30
+ ```python
31
+ import numpy as np
32
+ from k3_node.layers import CrystalConv
33
+
34
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
35
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
36
+ edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
37
+
38
+ layer = CrystalConv(channels=8, edge_dim=3)
39
+ out = layer(x, edge_index, edge_attr)
40
+ print(tuple(out.shape)) # (10, 8)
41
+ ```
42
+ """
43
+ def __init__(
44
+ self,
45
+ channels=None,
46
+ edge_dim=None,
47
+ aggregate="sum",
48
+ activation=None,
49
+ use_bias=True,
50
+ kernel_initializer="glorot_uniform",
51
+ bias_initializer="zeros",
52
+ kernel_regularizer=None,
53
+ bias_regularizer=None,
54
+ activity_regularizer=None,
55
+ kernel_constraint=None,
56
+ bias_constraint=None,
57
+ **kwargs,
58
+ ):
59
+
60
+ super().__init__(
61
+ aggregate=aggregate,
62
+ activation=activation,
63
+ use_bias=use_bias,
64
+ kernel_initializer=kernel_initializer,
65
+ bias_initializer=bias_initializer,
66
+ kernel_regularizer=kernel_regularizer,
67
+ bias_regularizer=bias_regularizer,
68
+ activity_regularizer=activity_regularizer,
69
+ kernel_constraint=kernel_constraint,
70
+ bias_constraint=bias_constraint,
71
+ **kwargs,
72
+ )
73
+ self.channels = channels
74
+ self.edge_dim = edge_dim
75
+
76
+ def build(self, input_shape=None):
77
+ layer_kwargs = dict(
78
+ kernel_initializer=self.kernel_initializer,
79
+ bias_initializer=self.bias_initializer,
80
+ kernel_regularizer=self.kernel_regularizer,
81
+ bias_regularizer=self.bias_regularizer,
82
+ kernel_constraint=self.kernel_constraint,
83
+ bias_constraint=self.bias_constraint,
84
+ dtype=self.dtype,
85
+ )
86
+ if self.channels is not None:
87
+ channels = self.channels
88
+ elif input_shape is not None:
89
+ if isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and isinstance(input_shape[0], (list, tuple)):
90
+ channels = input_shape[0][-1]
91
+ elif isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and isinstance(input_shape[-1], int):
92
+ channels = input_shape[-1]
93
+ else:
94
+ channels = 16
95
+ else:
96
+ channels = 16
97
+
98
+ self.channels = channels
99
+ self.dense_f = Dense(channels, activation="sigmoid", **layer_kwargs)
100
+ self.dense_s = Dense(channels, activation=self.activation, **layer_kwargs)
101
+
102
+ self.built = True
103
+
104
+ def call(self, x, edge_index=None, edge_attr=None, **kwargs):
105
+ if not self.built:
106
+ x_shape = getattr(x, "shape", None)
107
+ self.build(x_shape)
108
+
109
+ if edge_index is None and isinstance(x, (tuple, list)):
110
+ x_in, a, e = self.get_inputs(x)
111
+ return self.propagate(x_in, a, e, **kwargs)
112
+
113
+ if edge_attr is None:
114
+ edge_attr = kwargs.get("e", None)
115
+
116
+ return self.propagate(edge_index, x=x, edge_attr=edge_attr)
117
+
118
+ def message(self, x_i=None, x_j=None, edge_attr=None, x=None, e=None, **kwargs):
119
+ if x_i is None and x is not None:
120
+ x_i = self.get_targets(x)
121
+ x_j = self.get_sources(x)
122
+ if edge_attr is None:
123
+ edge_attr = e
124
+
125
+ to_concat = [x_i, x_j]
126
+ if e is not None:
127
+ to_concat += [e]
128
+ if edge_attr is not None:
129
+ to_concat.append(edge_attr)
130
+ z = ops.concatenate(to_concat, axis=-1)
131
+ output = self.dense_s(z) * self.dense_f(z)
132
+
133
+ return output
134
+
135
+ def update(self, embeddings, x=None, **kwargs):
136
+ if x is None:
137
+ return embeddings
138
+ if isinstance(x, (tuple, list)):
139
+ x = x[1] if x[1] is not None else x[0]
140
+ return x + embeddings
@@ -0,0 +1,84 @@
1
+ from k3_node.layers.conv.gat_conv import GATConv
2
+ from k3_node.layers.conv.sage_conv import SAGEConv
3
+ from k3_node.layers.conv.rgcn_conv import CuGraphRGCNConv
4
+
5
+
6
+ class CuGraphGATConv(GATConv):
7
+ r"""An optimized / multi-backend compatible version of :class:`GATConv`.
8
+
9
+ Example:
10
+ ```python
11
+ import numpy as np
12
+ from k3_node.layers import CuGraphGATConv
13
+
14
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
15
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
16
+
17
+ layer = CuGraphGATConv(8, 16, heads=2)
18
+ out = layer(x, edge_index)
19
+ print(tuple(out.shape)) # (10, 32)
20
+ ```
21
+ """
22
+
23
+ def __init__(
24
+ self,
25
+ in_channels: int,
26
+ out_channels: int,
27
+ heads: int = 1,
28
+ concat: bool = True,
29
+ negative_slope: float = 0.2,
30
+ bias: bool = True,
31
+ **kwargs,
32
+ ):
33
+ super().__init__(
34
+ in_channels=in_channels,
35
+ out_channels=out_channels,
36
+ heads=heads,
37
+ concat=concat,
38
+ negative_slope=negative_slope,
39
+ bias=bias,
40
+ **kwargs,
41
+ )
42
+
43
+
44
+ class CuGraphSAGEConv(SAGEConv):
45
+ r"""An optimized / multi-backend compatible version of :class:`SAGEConv`.
46
+
47
+ Example:
48
+ ```python
49
+ import numpy as np
50
+ from k3_node.layers import CuGraphSAGEConv
51
+
52
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
53
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
54
+
55
+ layer = CuGraphSAGEConv(8, 16)
56
+ out = layer(x, edge_index)
57
+ print(tuple(out.shape)) # (10, 16)
58
+ ```
59
+ """
60
+
61
+ def __init__(
62
+ self,
63
+ in_channels: int,
64
+ out_channels: int,
65
+ aggr: str = "mean",
66
+ normalize: bool = False,
67
+ root_weight: bool = True,
68
+ project: bool = False,
69
+ bias: bool = True,
70
+ **kwargs,
71
+ ):
72
+ super().__init__(
73
+ in_channels=in_channels,
74
+ out_channels=out_channels,
75
+ aggr=aggr,
76
+ normalize=normalize,
77
+ root_weight=root_weight,
78
+ project=project,
79
+ bias=bias,
80
+ **kwargs,
81
+ )
82
+
83
+
84
+ __all__ = ["CuGraphGATConv", "CuGraphSAGEConv", "CuGraphRGCNConv"]
@@ -0,0 +1,144 @@
1
+ from keras import layers, ops
2
+
3
+ from k3_node.layers.conv.conv import Conv
4
+ from k3_node.ops import normalized_adjacency, polyval
5
+
6
+
7
+ class DiffuseFeatures(layers.Layer):
8
+ def __init__(
9
+ self,
10
+ num_diffusion_steps,
11
+ kernel_initializer,
12
+ kernel_regularizer,
13
+ kernel_constraint,
14
+ **kwargs,
15
+ ):
16
+ super().__init__(**kwargs)
17
+
18
+ self.K = num_diffusion_steps
19
+ self.kernel_initializer = kernel_initializer
20
+ self.kernel_regularizer = kernel_regularizer
21
+ self.kernel_constraint = kernel_constraint
22
+
23
+ def build(self, input_shape):
24
+ # Initializing the kernel vector (R^K) (theta in paper)
25
+ self.kernel = self.add_weight(
26
+ shape=(self.K,),
27
+ name="kernel",
28
+ initializer=self.kernel_initializer,
29
+ regularizer=self.kernel_regularizer,
30
+ constraint=self.kernel_constraint,
31
+ )
32
+
33
+ def call(self, inputs):
34
+ x, a = inputs
35
+
36
+ diffusion_matrix = polyval(self.kernel, a)
37
+ diffused_features = ops.matmul(diffusion_matrix, x)
38
+ H = ops.sum(diffused_features, axis=-1)
39
+ return ops.expand_dims(H, -1)
40
+
41
+
42
+ class DiffusionConv(Conv):
43
+ """
44
+ `k3_node.layers.DiffusionConv`
45
+ Implementation of Diffusion Convolutional Neural Networks (DCNN) layer
46
+
47
+ Args:
48
+ channels: The number of output channels.
49
+ K: The number of diffusion steps.
50
+ activation: Activation function to use.
51
+ kernel_initializer: Initializer for the `kernel` weights matrix.
52
+ kernel_regularizer: Regularizer for the `kernel` weights matrix.
53
+ kernel_constraint: Constraint for the `kernel` weights matrix.
54
+ **kwargs: Additional arguments to pass to the `Conv` superclass.
55
+
56
+ Example:
57
+ ```python
58
+ import numpy as np
59
+ from k3_node.layers import DiffusionConv
60
+
61
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
62
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
63
+
64
+ layer = DiffusionConv(channels=8, out_channels=16, K=2)
65
+ out = layer(x, edge_index)
66
+ print(tuple(out.shape)) # (10, 16)
67
+ ```
68
+ """
69
+ def __init__(
70
+ self,
71
+ channels,
72
+ out_channels=None,
73
+ K=6,
74
+ activation="tanh",
75
+ kernel_initializer="glorot_uniform",
76
+ kernel_regularizer=None,
77
+ kernel_constraint=None,
78
+ **kwargs,
79
+ ):
80
+ if out_channels is not None:
81
+ self.in_channels = channels
82
+ channels = out_channels
83
+ super().__init__(
84
+ activation=activation,
85
+ kernel_initializer=kernel_initializer,
86
+ kernel_regularizer=kernel_regularizer,
87
+ kernel_constraint=kernel_constraint,
88
+ **kwargs,
89
+ )
90
+
91
+ self.channels = channels
92
+ self.K = K + 1
93
+
94
+ def build(self, input_shape):
95
+ self.filters = [
96
+ DiffuseFeatures(
97
+ num_diffusion_steps=self.K,
98
+ kernel_initializer=self.kernel_initializer,
99
+ kernel_regularizer=self.kernel_regularizer,
100
+ kernel_constraint=self.kernel_constraint,
101
+ )
102
+ for _ in range(self.channels)
103
+ ]
104
+ for f in self.filters:
105
+ f.build(None)
106
+ super().build(input_shape)
107
+
108
+ def apply_filters(self, x, a):
109
+ diffused_features = []
110
+
111
+ for diffusion in self.filters:
112
+ diffused_feature = diffusion((x, a))
113
+ diffused_features.append(diffused_feature)
114
+
115
+ return ops.concatenate(diffused_features, -1)
116
+
117
+ def call(self, inputs, a=None, **kwargs):
118
+ if a is not None:
119
+ x = inputs
120
+ elif isinstance(inputs, (list, tuple)):
121
+ x, a = inputs
122
+ else:
123
+ x, a = inputs, None
124
+
125
+ if a is not None and hasattr(a, "shape") and len(a.shape) == 2 and a.shape[0] == 2 and a.shape[1] != 2:
126
+ num_nodes = ops.shape(x)[0]
127
+ a_dense = ops.zeros((num_nodes, num_nodes), dtype=x.dtype)
128
+ indices = ops.transpose(a, axes=[1, 0])
129
+ updates = ops.ones(shape=(ops.shape(a)[1],), dtype=x.dtype)
130
+ a = ops.scatter_update(a_dense, indices, updates)
131
+
132
+ output = self.apply_filters(x, a)
133
+
134
+ output = self.activation(output)
135
+
136
+ return output
137
+
138
+ @property
139
+ def config(self):
140
+ return {"channels": self.channels, "K": self.K - 1}
141
+
142
+ @staticmethod
143
+ def preprocess(a):
144
+ return normalized_adjacency(a)
@@ -0,0 +1,93 @@
1
+ import copy
2
+ from keras import ops
3
+ from keras.layers import Layer, Dense
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+
7
+
8
+ class DirGNNConv(Layer):
9
+ r"""A directed graph neural network operator from the
10
+ `"Directed Graph Neural Networks" <https://arxiv.org/abs/2301.07663>`_ paper.
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import DirGNNConv, GCNConv
16
+
17
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
18
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
19
+
20
+ # Runs the wrapped convolution on incoming and outgoing edges separately
21
+ layer = DirGNNConv(GCNConv(in_channels=8, out_channels=16))
22
+ out = layer(x, edge_index)
23
+ print(tuple(out.shape)) # (10, 16)
24
+ ```
25
+ """
26
+ def __init__(
27
+ self,
28
+ conv: MessagePassing,
29
+ alpha: float = 0.5,
30
+ root_weight: bool = True,
31
+ **kwargs,
32
+ ):
33
+ super().__init__(**kwargs)
34
+
35
+ self.alpha = alpha
36
+ self.root_weight = root_weight
37
+ self.conv = conv
38
+
39
+ try:
40
+ self.conv_in = conv.__class__.from_config(conv.get_config())
41
+ self.conv_out = conv.__class__.from_config(conv.get_config())
42
+ except Exception:
43
+ self.conv_in = copy.deepcopy(conv)
44
+ self.conv_out = copy.deepcopy(conv)
45
+
46
+ if hasattr(self.conv_in, "add_self_loops"):
47
+ self.conv_in.add_self_loops = False
48
+ self.conv_out.add_self_loops = False
49
+ if hasattr(self.conv_in, "root_weight"):
50
+ self.conv_in.root_weight = False
51
+ self.conv_out.root_weight = False
52
+
53
+ in_channels = getattr(conv, "in_channels", None)
54
+ out_channels = getattr(conv, "out_channels", None)
55
+ if isinstance(in_channels, (list, tuple)):
56
+ in_channels = in_channels[0]
57
+ self.in_channels = in_channels
58
+ self.out_channels = out_channels
59
+
60
+ if root_weight and out_channels is not None:
61
+ self.lin = Dense(out_channels, use_bias=True)
62
+ else:
63
+ self.lin = None
64
+
65
+ def build(self, input_shape=None):
66
+ if self.in_channels is not None:
67
+ if hasattr(self.conv_in, "build"):
68
+ self.conv_in.build((None, self.in_channels))
69
+ if hasattr(self.conv_out, "build"):
70
+ self.conv_out.build((None, self.in_channels))
71
+ if self.lin is not None:
72
+ self.lin.build((None, self.in_channels))
73
+ self.built = True
74
+
75
+ def call(self, inputs, edge_index=None, **kwargs):
76
+ if edge_index is None:
77
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
78
+ x, edge_index = inputs
79
+ else:
80
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
81
+ else:
82
+ x = inputs
83
+
84
+ x_in = self.conv_in(x, edge_index)
85
+ edge_index_rev = ops.stack([edge_index[1], edge_index[0]], axis=0)
86
+ x_out = self.conv_out(x, edge_index_rev)
87
+
88
+ out = self.alpha * x_out + (1.0 - self.alpha) * x_in
89
+
90
+ if self.lin is not None:
91
+ out = out + self.lin(x)
92
+
93
+ return out
@@ -0,0 +1,192 @@
1
+ import math
2
+ import keras
3
+ from keras import ops
4
+ from k3_node.layers.conv.gcn_conv import gcn_norm
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+
7
+
8
+ def restricted_softmax(src, axis: int = -1, margin: float = 0.0):
9
+ src_max = ops.maximum(ops.max(src, axis=axis, keepdims=True), 0.0)
10
+ out = ops.exp(src - src_max)
11
+ denom = ops.sum(out, axis=axis, keepdims=True) + ops.exp(margin - src_max)
12
+ return out / denom
13
+
14
+
15
+ class GroupedLinear(keras.layers.Layer):
16
+ def __init__(self, in_channels: int, out_channels: int, groups: int = 1, bias: bool = True, **kwargs):
17
+ super().__init__(**kwargs)
18
+ assert in_channels % groups == 0 and out_channels % groups == 0
19
+ self.in_channels = in_channels
20
+ self.out_channels = out_channels
21
+ self.groups = groups
22
+ self.use_bias = bias
23
+
24
+ def build(self, input_shape=None):
25
+ self.weight = self.add_weight(
26
+ shape=(self.groups, self.in_channels // self.groups, self.out_channels // self.groups),
27
+ initializer="glorot_uniform",
28
+ trainable=True,
29
+ name="weight",
30
+ )
31
+ if self.use_bias:
32
+ self.bias = self.add_weight(
33
+ shape=(self.out_channels,),
34
+ initializer="zeros",
35
+ trainable=True,
36
+ name="bias",
37
+ )
38
+ else:
39
+ self.bias = None
40
+ super().build(input_shape)
41
+
42
+ def call(self, src):
43
+ if self.groups > 1:
44
+ orig_shape = ops.shape(src)
45
+ # Flatten batch dims
46
+ src_flat = ops.reshape(src, (-1, self.groups, self.in_channels // self.groups))
47
+ src_trans = ops.transpose(src_flat, (1, 0, 2))
48
+ out = ops.matmul(src_trans, self.weight)
49
+ out = ops.transpose(out, (1, 0, 2))
50
+ out_shape = tuple(orig_shape[:-1]) + (self.out_channels,)
51
+ out = ops.reshape(out, out_shape)
52
+ else:
53
+ out = ops.matmul(src, self.weight[0])
54
+
55
+ if self.bias is not None:
56
+ out = out + self.bias
57
+ return out
58
+
59
+
60
+ class DNAMultiHead(keras.layers.Layer):
61
+ def __init__(self, in_channels: int, out_channels: int, heads: int = 1, groups: int = 1, bias: bool = True, **kwargs):
62
+ super().__init__(**kwargs)
63
+ self.in_channels = in_channels
64
+ self.out_channels = out_channels
65
+ self.heads = heads
66
+ self.groups = groups
67
+ self.use_bias = bias
68
+
69
+ self.lin_q = GroupedLinear(in_channels, out_channels, groups, bias)
70
+ self.lin_k = GroupedLinear(in_channels, out_channels, groups, bias)
71
+ self.lin_v = GroupedLinear(in_channels, out_channels, groups, bias)
72
+
73
+ def build(self, input_shape=None):
74
+ self.lin_q.build()
75
+ self.lin_k.build()
76
+ self.lin_v.build()
77
+ super().build(input_shape)
78
+
79
+ def call(self, query, key, value):
80
+ q = self.lin_q(query)
81
+ k = self.lin_k(key)
82
+ v = self.lin_v(value)
83
+
84
+ # q: (E, 1, C) -> (E, heads, 1, C // heads)
85
+ E = ops.shape(q)[0]
86
+ H = self.heads
87
+ D = self.out_channels // H
88
+
89
+ q_entries = ops.shape(q)[1]
90
+ k_entries = ops.shape(k)[1]
91
+
92
+ q = ops.transpose(ops.reshape(q, (E, q_entries, H, D)), (0, 2, 1, 3))
93
+ k = ops.transpose(ops.reshape(k, (E, k_entries, H, D)), (0, 2, 1, 3))
94
+ v = ops.transpose(ops.reshape(v, (E, k_entries, H, D)), (0, 2, 1, 3))
95
+
96
+ # score: (E, H, q_entries, k_entries)
97
+ score = ops.matmul(q, ops.transpose(k, (0, 1, 3, 2))) / math.sqrt(D)
98
+ score = restricted_softmax(score, axis=-1)
99
+
100
+ out = ops.matmul(score, v) # (E, H, q_entries, D)
101
+ out = ops.transpose(out, (0, 2, 1, 3)) # (E, q_entries, H, D)
102
+ return ops.reshape(out, (E, q_entries, self.out_channels))
103
+
104
+
105
+ class DNAConv(MessagePassing):
106
+ r"""The dynamic neighborhood aggregation operator from the `"Just Jump:
107
+ Towards Dynamic Neighborhood Aggregation in Graph Neural Networks"
108
+ <https://arxiv.org/abs/1904.04849>`_ paper.
109
+
110
+ Args:
111
+ channels (int): Size of each input/output sample.
112
+ heads (int, optional): Number of multi-head-attentions. (default: :obj:`1`)
113
+ groups (int, optional): Number of groups for linear projections. (default: :obj:`1`)
114
+ dropout (float, optional): Dropout probability. (default: :obj:`0.0`)
115
+ cached (bool, optional): Whether to cache GCN normalization. (default: :obj:`False`)
116
+ normalize (bool, optional): Whether to apply symmetric normalization. (default: :obj:`True`)
117
+ add_self_loops (bool, optional): Whether to add self-loops. (default: :obj:`True`)
118
+ bias (bool, optional): Whether to learn an additive bias. (default: :obj:`True`)
119
+
120
+ Example:
121
+ ```python
122
+ import numpy as np
123
+ from k3_node.layers import DNAConv
124
+
125
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
126
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
127
+
128
+ # DNAConv attends over the representations of all previous layers: [num_nodes, num_layers, channels]
129
+ x_all = np.random.rand(10, 3, 8).astype("float32")
130
+ layer = DNAConv(channels=8, heads=2, groups=2)
131
+ out = layer(x_all, edge_index)
132
+ print(tuple(out.shape)) # (10, 8)
133
+ ```
134
+ """
135
+
136
+ def __init__(
137
+ self,
138
+ channels: int,
139
+ heads: int = 1,
140
+ groups: int = 1,
141
+ dropout: float = 0.0,
142
+ cached: bool = False,
143
+ normalize: bool = True,
144
+ add_self_loops: bool = True,
145
+ bias: bool = True,
146
+ **kwargs,
147
+ ):
148
+ kwargs.setdefault("aggr", "add")
149
+ super().__init__(node_dim=0, **kwargs)
150
+
151
+ self.channels = channels
152
+ self.heads = heads
153
+ self.groups = groups
154
+ self.dropout = dropout
155
+ self.cached = cached
156
+ self.normalize = normalize
157
+ self.add_self_loops = add_self_loops
158
+ self.use_bias = bias
159
+
160
+ self.multi_head = DNAMultiHead(channels, channels, heads, groups, bias)
161
+
162
+ def build(self, input_shape=None):
163
+ self.multi_head.build()
164
+ super().build(input_shape)
165
+
166
+ def call(self, x, edge_index, edge_weight=None):
167
+ if not self.built:
168
+ self.build()
169
+
170
+ if len(ops.shape(x)) == 2:
171
+ x = ops.expand_dims(x, axis=1)
172
+
173
+ num_nodes = ops.shape(x)[0]
174
+ if self.normalize:
175
+ edge_index, edge_weight = gcn_norm(
176
+ edge_index,
177
+ edge_weight,
178
+ num_nodes=num_nodes,
179
+ add_self_loops=self.add_self_loops,
180
+ dtype=x.dtype,
181
+ )
182
+
183
+ return self.propagate(edge_index, x=x, edge_weight=edge_weight)
184
+
185
+ def message(self, x_i, x_j, edge_weight=None):
186
+ q = x_i[:, -1:] # (E, 1, C)
187
+ out = self.multi_head(q, x_j, x_j) # (E, 1, C)
188
+ out = ops.squeeze(out, axis=1) # (E, C)
189
+ if edge_weight is not None:
190
+ out = ops.expand_dims(edge_weight, axis=-1) * out
191
+ return out
192
+