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,88 @@
1
+ from keras import initializers, layers, ops
2
+
3
+
4
+ class DenseGINConv(layers.Layer):
5
+ r"""Applies the dense Graph Isomorphism Network (GIN) convolutional operator
6
+ from the `"How Powerful are Graph Neural Networks?"
7
+ <https://arxiv.org/abs/1810.00826>`_ paper.
8
+
9
+ .. math::
10
+ \mathbf{X}^{\prime} = \mathrm{MLP} \left( \left( \mathbf{A} + (1 + \epsilon)
11
+ \cdot \mathbf{I} \right) \cdot \mathbf{X} \right)
12
+
13
+ Args:
14
+ nn (callable): A neural network layer (e.g. MLP or Dense) mapping feature
15
+ representations to output representations.
16
+ eps (float, optional): (Initial) :math:`\epsilon`-value. (default: :obj:`0.0`)
17
+ train_eps (bool, optional): If set to :obj:`True`, :math:`\epsilon` will
18
+ be a trainable parameter. (default: :obj:`False`)
19
+
20
+ Example:
21
+ ```python
22
+ import numpy as np
23
+ import keras
24
+ from k3_node.layers import DenseGINConv
25
+
26
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
27
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
28
+
29
+ layer = DenseGINConv(keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)]))
30
+ out = layer(x, adj)
31
+ print(tuple(out.shape)) # (2, 10, 16)
32
+ ```
33
+ """
34
+ def __init__(
35
+ self,
36
+ nn,
37
+ eps: float = 0.0,
38
+ train_eps: bool = False,
39
+ **kwargs
40
+ ):
41
+ super().__init__(**kwargs)
42
+ self.nn = nn
43
+ self.initial_eps = eps
44
+ self.train_eps = train_eps
45
+
46
+ self.eps = self.add_weight(
47
+ shape=(1,),
48
+ initializer=initializers.Constant(eps),
49
+ trainable=train_eps,
50
+ name="eps",
51
+ )
52
+
53
+ def reset_parameters(self):
54
+ if hasattr(self.nn, 'reset_parameters'):
55
+ self.nn.reset_parameters()
56
+ self.eps.assign(ops.cast(ops.convert_to_tensor([self.initial_eps]), dtype=self.eps.dtype))
57
+
58
+ def call(self, x, adj, mask=None, add_loop: bool = True):
59
+ is_2d_x = (len(ops.shape(x)) == 2)
60
+ if is_2d_x:
61
+ x = ops.expand_dims(x, axis=0)
62
+ if len(ops.shape(adj)) == 2:
63
+ adj = ops.expand_dims(adj, axis=0)
64
+
65
+ N = ops.shape(adj)[1]
66
+
67
+ out = ops.matmul(adj, x)
68
+ if add_loop:
69
+ out = (1.0 + self.eps) * x + out
70
+
71
+ out = self.nn(out)
72
+
73
+ if mask is not None:
74
+ out = out * ops.cast(ops.reshape(mask, (-1, N, 1)), x.dtype)
75
+
76
+ if is_2d_x and ops.shape(out)[0] == 1:
77
+ out = ops.squeeze(out, axis=0)
78
+
79
+ return out
80
+
81
+ def compute_output_shape(self, input_shape):
82
+ if hasattr(self.nn, 'compute_output_shape'):
83
+ return self.nn.compute_output_shape(input_shape)
84
+ return input_shape
85
+
86
+ def __repr__(self) -> str:
87
+ return f'{self.__class__.__name__}(nn={self.nn})'
88
+
@@ -0,0 +1,95 @@
1
+ from typing import Optional
2
+ from keras import layers, ops
3
+ from .linear import Linear
4
+
5
+
6
+ class DenseGraphConv(layers.Layer):
7
+ r"""See :class:`torch_geometric.nn.conv.GraphConv`.
8
+
9
+ Example:
10
+ ```python
11
+ import numpy as np
12
+ from k3_node.layers import DenseGraphConv
13
+
14
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
15
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
16
+
17
+ layer = DenseGraphConv(in_channels=8, out_channels=16)
18
+ out = layer(x, adj)
19
+ print(tuple(out.shape)) # (2, 10, 16)
20
+ ```
21
+ """
22
+ def __init__(
23
+ self,
24
+ in_channels: int,
25
+ out_channels: int,
26
+ aggr: str = 'add',
27
+ bias: bool = True,
28
+ **kwargs
29
+ ):
30
+ if aggr not in ('add', 'mean', 'max'):
31
+ raise ValueError(f"aggr must be one of 'add', 'mean', 'max' (got '{aggr}')")
32
+ super().__init__(**kwargs)
33
+
34
+ self.in_channels = in_channels
35
+ self.out_channels = out_channels
36
+ self.aggr = aggr
37
+ self.use_bias = bias
38
+
39
+ self.lin_rel = Linear(in_channels, out_channels, bias=bias, name="lin_rel")
40
+ self.lin_root = Linear(in_channels, out_channels, bias=False, name="lin_root")
41
+
42
+ def reset_parameters(self):
43
+ self.lin_rel.reset_parameters()
44
+ self.lin_root.reset_parameters()
45
+
46
+ def build(self, input_shape):
47
+ if isinstance(input_shape, (tuple, list)):
48
+ x_shape = input_shape[0]
49
+ else:
50
+ x_shape = input_shape
51
+ self.lin_rel.build(x_shape)
52
+ self.lin_root.build(x_shape)
53
+ super().build(input_shape)
54
+
55
+ def call(self, x, adj, mask=None):
56
+ is_2d_x = (len(ops.shape(x)) == 2)
57
+ if is_2d_x:
58
+ x = ops.expand_dims(x, axis=0)
59
+ if len(ops.shape(adj)) == 2:
60
+ adj = ops.expand_dims(adj, axis=0)
61
+
62
+ N = ops.shape(x)[1]
63
+
64
+ if self.aggr == 'add':
65
+ out = ops.matmul(adj, x)
66
+ elif self.aggr == 'mean':
67
+ deg = ops.maximum(ops.sum(adj, axis=-1, keepdims=True), 1.0)
68
+ out = ops.matmul(adj, x) / deg
69
+ elif self.aggr == 'max':
70
+ x_src = ops.expand_dims(x, axis=1) # [B, 1, N, C]
71
+ adj_expanded = ops.expand_dims(adj, axis=-1) # [B, N, N, 1]
72
+ masked = ops.where(adj_expanded != 0, x_src, -1e9)
73
+ out = ops.max(masked, axis=2) # [B, N, C]
74
+ out = ops.where(out <= -1e8, 0.0, out)
75
+
76
+ out = self.lin_rel(out) + self.lin_root(x)
77
+
78
+ if mask is not None:
79
+ out = out * ops.cast(ops.reshape(mask, (-1, N, 1)), x.dtype)
80
+
81
+ if is_2d_x and ops.shape(out)[0] == 1:
82
+ out = ops.squeeze(out, axis=0)
83
+
84
+ return out
85
+
86
+ def compute_output_shape(self, input_shape):
87
+ if isinstance(input_shape, (tuple, list)):
88
+ x_shape = input_shape[0]
89
+ else:
90
+ x_shape = input_shape
91
+ return (*x_shape[:-1], self.out_channels)
92
+
93
+ def __repr__(self) -> str:
94
+ return (f'{self.__class__.__name__}({self.in_channels}, '
95
+ f'{self.out_channels})')
@@ -0,0 +1,85 @@
1
+ from keras import layers, ops
2
+ from .linear import Linear
3
+
4
+
5
+ class DenseSAGEConv(layers.Layer):
6
+ r"""See :class:`torch_geometric.nn.conv.SAGEConv`.
7
+
8
+ Example:
9
+ ```python
10
+ import numpy as np
11
+ from k3_node.layers import DenseSAGEConv
12
+
13
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
14
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
15
+
16
+ layer = DenseSAGEConv(in_channels=8, out_channels=16)
17
+ out = layer(x, adj)
18
+ print(tuple(out.shape)) # (2, 10, 16)
19
+ ```
20
+ """
21
+ def __init__(
22
+ self,
23
+ in_channels: int,
24
+ out_channels: int,
25
+ normalize: bool = False,
26
+ bias: bool = True,
27
+ **kwargs
28
+ ):
29
+ super().__init__(**kwargs)
30
+ self.in_channels = in_channels
31
+ self.out_channels = out_channels
32
+ self.normalize = normalize
33
+ self.use_bias = bias
34
+
35
+ self.lin_rel = Linear(in_channels, out_channels, bias=False, name="lin_rel")
36
+ self.lin_root = Linear(in_channels, out_channels, bias=bias, name="lin_root")
37
+
38
+ def reset_parameters(self):
39
+ self.lin_rel.reset_parameters()
40
+ self.lin_root.reset_parameters()
41
+
42
+ def build(self, input_shape):
43
+ if isinstance(input_shape, (tuple, list)):
44
+ x_shape = input_shape[0]
45
+ else:
46
+ x_shape = input_shape
47
+ self.lin_rel.build(x_shape)
48
+ self.lin_root.build(x_shape)
49
+ super().build(input_shape)
50
+
51
+ def call(self, x, adj, mask=None):
52
+ is_2d_x = (len(ops.shape(x)) == 2)
53
+ if is_2d_x:
54
+ x = ops.expand_dims(x, axis=0)
55
+ if len(ops.shape(adj)) == 2:
56
+ adj = ops.expand_dims(adj, axis=0)
57
+
58
+ N = ops.shape(adj)[1]
59
+
60
+ deg = ops.maximum(ops.sum(adj, axis=-1, keepdims=True), 1.0)
61
+ out = ops.matmul(adj, x) / deg
62
+ out = self.lin_rel(out) + self.lin_root(x)
63
+
64
+ if self.normalize:
65
+ norm = ops.maximum(ops.sqrt(ops.sum(ops.power(out, 2), axis=-1, keepdims=True)), 1e-12)
66
+ out = out / norm
67
+
68
+ if mask is not None:
69
+ out = out * ops.cast(ops.reshape(mask, (-1, N, 1)), x.dtype)
70
+
71
+ if is_2d_x and ops.shape(out)[0] == 1:
72
+ out = ops.squeeze(out, axis=0)
73
+
74
+ return out
75
+
76
+ def compute_output_shape(self, input_shape):
77
+ if isinstance(input_shape, (tuple, list)):
78
+ x_shape = input_shape[0]
79
+ else:
80
+ x_shape = input_shape
81
+ return (*x_shape[:-1], self.out_channels)
82
+
83
+ def __repr__(self) -> str:
84
+ return (f'{self.__class__.__name__}({self.in_channels}, '
85
+ f'{self.out_channels})')
@@ -0,0 +1,76 @@
1
+ from typing import Optional, Tuple
2
+ from keras import ops
3
+
4
+
5
+ def dense_diff_pool(
6
+ x,
7
+ adj,
8
+ s,
9
+ mask: Optional[any] = None,
10
+ normalize: bool = True,
11
+ ) -> Tuple[any, any, any, any]:
12
+ r"""The differentiable pooling operator from the `"Hierarchical Graph
13
+ Representation Learning with Differentiable Pooling"
14
+ <https://arxiv.org/abs/1806.08804>`_ paper.
15
+
16
+ .. math::
17
+ \mathbf{X}^{\prime} &= {\mathrm{softmax}(\mathbf{S})}^{\top} \cdot
18
+ \mathbf{X}
19
+
20
+ \mathbf{A}^{\prime} &= {\mathrm{softmax}(\mathbf{S})}^{\top} \cdot
21
+ \mathbf{A} \cdot \mathrm{softmax}(\mathbf{S})
22
+
23
+ Args:
24
+ x: Node feature tensor [B, N, F] or [N, F].
25
+ adj: Adjacency tensor [B, N, N] or [N, N].
26
+ s: Assignment tensor [B, N, C] or [N, C].
27
+ mask: Mask tensor [B, N] indicating valid nodes. (default: None)
28
+ normalize: If set to False, link prediction loss is not divided by total elements.
29
+
30
+ Example:
31
+ ```python
32
+ import numpy as np
33
+ from k3_node.layers import dense_diff_pool
34
+
35
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
36
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
37
+ s = np.random.rand(2, 10, 3).astype("float32") # assignment scores for 3 clusters
38
+
39
+ x_pool, adj_pool, link_loss, entropy_loss = dense_diff_pool(x, adj, s)
40
+ print(tuple(x_pool.shape), tuple(adj_pool.shape)) # (2, 3, 8) (2, 3, 3)
41
+ ```
42
+ """
43
+ # Plain NumPy inputs cannot be mixed with backend tensors (e.g. `ndarray @ torch.Tensor`).
44
+ x, adj, s = ops.convert_to_tensor(x), ops.convert_to_tensor(adj), ops.convert_to_tensor(s)
45
+ if len(ops.shape(x)) == 2:
46
+ x = ops.expand_dims(x, axis=0)
47
+ if len(ops.shape(adj)) == 2:
48
+ adj = ops.expand_dims(adj, axis=0)
49
+ if len(ops.shape(s)) == 2:
50
+ s = ops.expand_dims(s, axis=0)
51
+
52
+ batch_size = ops.shape(x)[0]
53
+ num_nodes = ops.shape(x)[1]
54
+
55
+ s = ops.softmax(s, axis=-1)
56
+
57
+ if mask is not None:
58
+ mask_m = ops.cast(ops.reshape(mask, (batch_size, num_nodes, 1)), x.dtype)
59
+ x = x * mask_m
60
+ s = s * mask_m
61
+
62
+ s_t = ops.transpose(s, (0, 2, 1))
63
+
64
+ out = ops.matmul(s_t, x)
65
+ out_adj = ops.matmul(ops.matmul(s_t, adj), s)
66
+
67
+ link = adj - ops.matmul(s, s_t)
68
+ link_loss = ops.sqrt(ops.sum(ops.power(link, 2)))
69
+ if normalize:
70
+ numel = ops.cast(ops.prod(ops.shape(adj)), link_loss.dtype)
71
+ link_loss = link_loss / numel
72
+
73
+ ent_loss = ops.mean(ops.sum(-s * ops.log(s + 1e-15), axis=-1))
74
+
75
+ return out, out_adj, link_loss, ent_loss
76
+
@@ -0,0 +1,223 @@
1
+ from typing import List, Optional, Tuple, Union
2
+ from keras import layers, ops
3
+
4
+ from .linear import Linear
5
+
6
+
7
+ class MLP(layers.Layer):
8
+ r"""A simple Multi-Layer Perceptron (MLP) model."""
9
+ def __init__(
10
+ self,
11
+ channel_list: List[int],
12
+ act: Optional[Union[str, any]] = None,
13
+ norm: Optional[Union[str, any]] = None,
14
+ **kwargs,
15
+ ):
16
+ super().__init__(**kwargs)
17
+ self.channel_list = channel_list
18
+ self.lins = [
19
+ Linear(in_c, out_c)
20
+ for in_c, out_c in zip(channel_list[:-1], channel_list[1:])
21
+ ]
22
+ self.act = layers.Activation(act) if act is not None else None
23
+
24
+ @property
25
+ def in_channels(self) -> int:
26
+ r"""Size of each input sample."""
27
+ return self.channel_list[0]
28
+
29
+ @property
30
+ def out_channels(self) -> int:
31
+ r"""Size of each output sample."""
32
+ return self.channel_list[-1]
33
+
34
+ def reset_parameters(self):
35
+ r"""Resets all learnable parameters of the module."""
36
+ for lin in self.lins:
37
+ if hasattr(lin, "reset_parameters"):
38
+ lin.reset_parameters()
39
+
40
+ def call(self, x):
41
+ for i, lin in enumerate(self.lins):
42
+ x = lin(x)
43
+ if self.act is not None and i < len(self.lins) - 1:
44
+ x = self.act(x)
45
+ return x
46
+
47
+
48
+ class DMoNPooling(layers.Layer):
49
+ r"""The spectral modularity pooling operator from the `"Graph Clustering
50
+ with Graph Neural Networks" <https://arxiv.org/abs/2006.16904>`_ paper.
51
+
52
+ .. math::
53
+ \mathbf{X}^{\prime} &= {\mathrm{softmax}(\mathbf{S})}^{\top} \cdot
54
+ \mathbf{X}
55
+
56
+ \mathbf{A}^{\prime} &= {\mathrm{softmax}(\mathbf{S})}^{\top} \cdot
57
+ \mathbf{A} \cdot \mathrm{softmax}(\mathbf{S})
58
+
59
+ Args:
60
+ channels (int or List[int]): Size of each input sample. If given as a
61
+ list, will construct an MLP based on the given feature sizes.
62
+ k (int): The number of clusters.
63
+ dropout (float, optional): Dropout probability. (default: :obj:`0.0`)
64
+
65
+ Example:
66
+ ```python
67
+ import numpy as np
68
+ from k3_node.layers import DMoNPooling
69
+
70
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
71
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
72
+
73
+ layer = DMoNPooling(channels=8, k=3) # 3 clusters
74
+ s, x_pool, adj_pool, spectral_loss, ortho_loss, cluster_loss = layer(x, adj)
75
+ print(tuple(s.shape)) # (2, 10, 3): soft cluster assignments
76
+ print(tuple(x_pool.shape), tuple(adj_pool.shape)) # (2, 3, 8) (2, 3, 3)
77
+ ```
78
+ """
79
+ def __init__(
80
+ self,
81
+ channels: Union[int, List[int]],
82
+ k: int,
83
+ dropout: float = 0.0,
84
+ **kwargs,
85
+ ):
86
+ super().__init__(**kwargs)
87
+
88
+ if isinstance(channels, int):
89
+ channels = [channels]
90
+
91
+ self.channels = channels
92
+ self.k = k
93
+ self.dropout = dropout
94
+ self.mlp = MLP(channels + [k], act=None, norm=None)
95
+ self.drop = layers.Dropout(dropout) if dropout > 0.0 else None
96
+
97
+ def reset_parameters(self):
98
+ r"""Resets all learnable parameters of the module."""
99
+ self.mlp.reset_parameters()
100
+
101
+ def call(
102
+ self,
103
+ x,
104
+ adj,
105
+ mask: Optional[any] = None,
106
+ training: bool = False,
107
+ ) -> Tuple[any, any, any, any, any, any]:
108
+ r"""Forward pass.
109
+
110
+ Args:
111
+ x: Node feature tensor [B, N, F] or [N, F].
112
+ adj: Adjacency tensor [B, N, N] or [N, N].
113
+ mask: Mask tensor [B, N] indicating valid nodes. (default: None)
114
+ training: Whether layer is in training mode.
115
+ """
116
+ if len(ops.shape(x)) == 2:
117
+ x = ops.expand_dims(x, axis=0)
118
+ if len(ops.shape(adj)) == 2:
119
+ adj = ops.expand_dims(adj, axis=0)
120
+
121
+ s = self.mlp(x)
122
+ if self.drop is not None:
123
+ s = self.drop(s, training=training)
124
+ s = ops.softmax(s, axis=-1)
125
+ if mask is not None:
126
+ batch_size = ops.shape(x)[0]
127
+ num_nodes = ops.shape(x)[1]
128
+ mask_m = ops.cast(ops.reshape(mask, (batch_size, num_nodes, 1)), x.dtype)
129
+ s = s * mask_m
130
+ out, out_adj, spectral_loss, ortho_loss, cluster_loss = dense_dmon_pool(x, adj, s, mask=mask)
131
+ return s, out, out_adj, spectral_loss, ortho_loss, cluster_loss
132
+
133
+ def __repr__(self) -> str:
134
+ return (f'{self.__class__.__name__}({self.mlp.in_channels}, '
135
+ f'num_clusters={self.mlp.out_channels})')
136
+
137
+
138
+ def dense_dmon_pool(
139
+ x,
140
+ adj,
141
+ s,
142
+ mask: Optional[any] = None,
143
+ ) -> Tuple[any, any, any, any, any]:
144
+ r"""Functional dense DMoN pooling operator.
145
+
146
+ Example:
147
+ ```python
148
+ import numpy as np
149
+ from k3_node.layers import dense_dmon_pool
150
+
151
+ x = np.random.rand(2, 10, 8).astype("float32") # batch of 2 graphs, 10 nodes, 8 features
152
+ adj = (np.random.rand(2, 10, 10) > 0.7).astype("float32") # dense adjacency matrices
153
+ s = np.random.rand(2, 10, 3).astype("float32") # assignment scores for 3 clusters
154
+
155
+ x_pool, adj_pool, spectral_loss, ortho_loss, cluster_loss = dense_dmon_pool(x, adj, s)
156
+ print(tuple(x_pool.shape), tuple(adj_pool.shape)) # (2, 3, 8) (2, 3, 3)
157
+ ```
158
+ """
159
+ # Plain NumPy inputs cannot be mixed with backend tensors (e.g. `ndarray @ torch.Tensor`).
160
+ x, adj, s = ops.convert_to_tensor(x), ops.convert_to_tensor(adj), ops.convert_to_tensor(s)
161
+ if len(ops.shape(x)) == 2:
162
+ x = ops.expand_dims(x, axis=0)
163
+ if len(ops.shape(adj)) == 2:
164
+ adj = ops.expand_dims(adj, axis=0)
165
+ if len(ops.shape(s)) == 2:
166
+ s = ops.expand_dims(s, axis=0)
167
+
168
+ batch_size = ops.shape(x)[0]
169
+ num_nodes = ops.shape(x)[1]
170
+ C = ops.shape(s)[-1]
171
+
172
+ if mask is None:
173
+ mask = ops.ones((batch_size, num_nodes, 1), dtype=x.dtype)
174
+ else:
175
+ mask = ops.cast(ops.reshape(mask, (batch_size, num_nodes, 1)), x.dtype)
176
+
177
+ x = x * mask
178
+ s = s * mask
179
+
180
+ out = ops.selu(ops.matmul(ops.swapaxes(s, 1, 2), x))
181
+ out_adj = ops.matmul(ops.matmul(ops.swapaxes(s, 1, 2), adj), s)
182
+
183
+ # Spectral loss:
184
+ degrees = ops.sum(adj, axis=-1, keepdims=True) * mask
185
+ degrees_t = ops.swapaxes(degrees, 1, 2)
186
+
187
+ m = ops.sum(degrees, axis=(1, 2)) / 2.0
188
+ m_expand = ops.broadcast_to(ops.reshape(m, (-1, 1, 1)), (batch_size, C, C))
189
+
190
+ ca = ops.matmul(ops.swapaxes(s, 1, 2), degrees)
191
+ cb = ops.matmul(degrees_t, s)
192
+
193
+ normalizer = ops.matmul(ca, cb) / 2.0 / (m_expand + 1e-15)
194
+ decompose = out_adj - normalizer
195
+ tr = ops.sum(ops.diagonal(decompose, axis1=1, axis2=2), axis=-1)
196
+ spectral_loss = ops.mean(-tr / 2.0 / (m + 1e-15))
197
+
198
+ # Orthogonality regularization:
199
+ ss = ops.matmul(ops.swapaxes(s, 1, 2), s)
200
+ i_s = ops.eye(C, dtype=ss.dtype)
201
+ norm_ss = ops.sqrt(ops.sum(ops.power(ss, 2), axis=(-1, -2), keepdims=True) + 1e-15)
202
+ norm_is = ops.sqrt(ops.cast(C, ss.dtype))
203
+ diff = (ss / norm_ss) - (ops.expand_dims(i_s, axis=0) / norm_is)
204
+ ortho_loss = ops.mean(ops.sqrt(ops.sum(ops.power(diff, 2), axis=(-1, -2)) + 1e-15))
205
+
206
+ # Cluster loss:
207
+ cluster_size = ops.sum(s, axis=1)
208
+ cluster_loss = ops.sqrt(ops.sum(ops.power(cluster_size, 2), axis=1) + 1e-15)
209
+ cluster_loss = ops.mean(cluster_loss / (ops.sum(mask, axis=1) + 1e-15) * norm_is - 1.0)
210
+
211
+ # Fix and normalize coarsened adjacency matrix:
212
+ eye_c = ops.expand_dims(ops.eye(C, dtype=out_adj.dtype), axis=0)
213
+ out_adj = out_adj * (1.0 - eye_c)
214
+ d = ops.sum(out_adj, axis=-1, keepdims=False)
215
+ d = ops.expand_dims(ops.sqrt(d), axis=1) + 1e-15
216
+ out_adj = (out_adj / d) / ops.swapaxes(d, 1, 2)
217
+
218
+ return out, out_adj, spectral_loss, ortho_loss, cluster_loss
219
+
220
+
221
+ dmon_pool = dense_dmon_pool
222
+
223
+