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,149 @@
1
+ from .message_passing import MessagePassing
2
+ from .simple_conv import SimpleConv
3
+ from .gcn_conv import GCNConv
4
+ from .cheb_conv import ChebConv
5
+ from .sage_conv import SAGEConv
6
+ from .graph_conv import GraphConv
7
+ from .gated_graph_conv import GatedGraphConv
8
+ from .res_gated_graph_conv import ResGatedGraphConv
9
+ from .gat_conv import GATConv, FusedGATConv
10
+ from .gatv2_conv import GATv2Conv
11
+ from .transformer_conv import TransformerConv
12
+ from .agnn_conv import AGNNConv
13
+ from .tag_conv import TAGConv
14
+ from .gin_conv import GINConv, GINEConv
15
+ from .arma_conv import ARMAConv
16
+ from .sg_conv import SGConv
17
+ from .ssg_conv import SSGConv
18
+ from .appnp import APPNP
19
+ from .appnp_conv import APPNPConv
20
+ from .mf_conv import MFConv
21
+ from .rgcn_conv import RGCNConv, FastRGCNConv, CuGraphRGCNConv
22
+ from .rgat_conv import RGATConv
23
+ from .signed_conv import SignedConv
24
+ from .dir_gnn_conv import DirGNNConv
25
+ from .antisymmetric_conv import AntiSymmetricConv
26
+ from .mixhop_conv import MixHopConv
27
+ from .pdn_conv import PDNConv
28
+ from .fa_conv import FAConv
29
+ from .film_conv import FiLMConv
30
+ from .supergat_conv import SuperGATConv
31
+ from .eg_conv import EGConv
32
+ from .pan_conv import PANConv
33
+ from .gen_conv import GENConv
34
+ from .pna_conv import PNAConv
35
+ from .le_conv import LEConv
36
+ from .cluster_gcn_conv import ClusterGCNConv
37
+ from .gcn2_conv import GCN2Conv
38
+ from .lg_conv import LGConv
39
+ from .nn_conv import NNConv
40
+ from .cg_conv import CGConv
41
+ from .edge_conv import EdgeConv, DynamicEdgeConv
42
+ from .general_conv import GeneralConv
43
+ from .point_conv import PointNetConv, PointConv
44
+ from .point_transformer_conv import PointTransformerConv
45
+ from .point_gnn_conv import PointGNNConv
46
+ from .ppf_conv import PPFConv
47
+ from .feast_conv import FeaStConv
48
+ from .gmm_conv import GMMConv
49
+ from .gravnet_conv import GravNetConv
50
+ from .meshcnn_conv import MeshCNNConv
51
+ from .x_conv import XConv
52
+ from .spline_conv import SplineConv
53
+ from .hetero_conv import HeteroConv
54
+ from .hgt_conv import HGTConv
55
+ from .han_conv import HANConv
56
+ from .heat_conv import HEATConv
57
+ from .hypergraph_conv import HypergraphConv
58
+ from .dna_conv import DNAConv
59
+ from .wl_conv import WLConv, WLConvContinuous
60
+ from .gps_conv import GPSConv
61
+ from .cugraph import CuGraphGATConv, CuGraphSAGEConv
62
+
63
+ ECConv = NNConv
64
+
65
+ # Legacy Spektral imports
66
+ from .crystal_conv import CrystalConv
67
+ from .diffusion_conv import DiffusionConv
68
+ from .gcn import GraphConvolution
69
+ from .graph_attention import GraphAttention
70
+ from .ppnp import PPNPPropagation
71
+
72
+ __all__ = [
73
+ "MessagePassing",
74
+ "SimpleConv",
75
+ "GCNConv",
76
+ "ChebConv",
77
+ "SAGEConv",
78
+ "GraphConv",
79
+ "GatedGraphConv",
80
+ "ResGatedGraphConv",
81
+ "GATConv",
82
+ "FusedGATConv",
83
+ "GATv2Conv",
84
+ "TransformerConv",
85
+ "AGNNConv",
86
+ "TAGConv",
87
+ "GINConv",
88
+ "GINEConv",
89
+ "ARMAConv",
90
+ "SGConv",
91
+ "SSGConv",
92
+ "APPNP",
93
+ "APPNPConv",
94
+ "MFConv",
95
+ "RGCNConv",
96
+ "FastRGCNConv",
97
+ "CuGraphRGCNConv",
98
+ "RGATConv",
99
+ "SignedConv",
100
+ "DirGNNConv",
101
+ "AntiSymmetricConv",
102
+ "MixHopConv",
103
+ "PDNConv",
104
+ "FAConv",
105
+ "FiLMConv",
106
+ "SuperGATConv",
107
+ "EGConv",
108
+ "PANConv",
109
+ "GENConv",
110
+ "PNAConv",
111
+ "LEConv",
112
+ "ClusterGCNConv",
113
+ "GCN2Conv",
114
+ "LGConv",
115
+ "NNConv",
116
+ "ECConv",
117
+ "CGConv",
118
+ "EdgeConv",
119
+ "DynamicEdgeConv",
120
+ "GeneralConv",
121
+ "PointNetConv",
122
+ "PointConv",
123
+ "PointTransformerConv",
124
+ "PointGNNConv",
125
+ "PPFConv",
126
+ "FeaStConv",
127
+ "GMMConv",
128
+ "GravNetConv",
129
+ "MeshCNNConv",
130
+ "XConv",
131
+ "SplineConv",
132
+ "HeteroConv",
133
+ "HGTConv",
134
+ "HANConv",
135
+ "HEATConv",
136
+ "HypergraphConv",
137
+ "DNAConv",
138
+ "WLConv",
139
+ "WLConvContinuous",
140
+ "GPSConv",
141
+ "CuGraphGATConv",
142
+ "CuGraphSAGEConv",
143
+ # Legacy Spektral
144
+ "CrystalConv",
145
+ "DiffusionConv",
146
+ "GraphConvolution",
147
+ "GraphAttention",
148
+ "PPNPPropagation",
149
+ ]
@@ -0,0 +1,120 @@
1
+ from typing import Optional
2
+ from keras import ops
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import (
6
+ add_self_loops,
7
+ extend_mask_for_self_loops,
8
+ mask_edge_logits,
9
+ remove_self_loops_masked,
10
+ softmax,
11
+ )
12
+
13
+
14
+ class AGNNConv(MessagePassing):
15
+ r"""The graph attentional propagation layer from the
16
+ `"Attention-based Graph Neural Network for Semi-Supervised Learning"
17
+ <https://arxiv.org/abs/1803.03735>`_ paper.
18
+
19
+ Example:
20
+ ```python
21
+ import numpy as np
22
+ from k3_node.layers import AGNNConv
23
+
24
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
25
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
26
+
27
+ layer = AGNNConv(requires_grad=True)
28
+ out = layer(x, edge_index)
29
+ print(tuple(out.shape)) # (10, 8)
30
+ ```
31
+ """
32
+ def __init__(
33
+ self,
34
+ requires_grad: bool = True,
35
+ add_self_loops: bool = True,
36
+ trainable: Optional[bool] = None,
37
+ aggregate: str = "add",
38
+ activation=None,
39
+ **kwargs,
40
+ ):
41
+ if trainable is not None:
42
+ requires_grad = trainable
43
+ kwargs.setdefault("aggr", aggregate)
44
+ super().__init__(activation=activation, **kwargs)
45
+
46
+ self.requires_grad = requires_grad
47
+ self.add_self_loops = add_self_loops
48
+
49
+ if requires_grad:
50
+ self.beta = self.add_weight(
51
+ shape=(1,),
52
+ initializer="ones",
53
+ name="beta",
54
+ )
55
+ else:
56
+ self.beta = None
57
+
58
+ def build(self, input_shape=None):
59
+ self.built = True
60
+
61
+ def call(self, inputs, edge_index=None, **kwargs):
62
+ if edge_index is None:
63
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
64
+ x, edge_index = inputs
65
+ else:
66
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
67
+ else:
68
+ x = inputs
69
+
70
+ if not self.built:
71
+ self.build()
72
+
73
+ # Check for legacy 2D matrix
74
+ is_legacy = False
75
+ if hasattr(edge_index, "shape") and len(edge_index.shape) == 2:
76
+ if (
77
+ edge_index.shape[0] is not None
78
+ and edge_index.shape[1] is not None
79
+ and edge_index.shape[0] != 2
80
+ and edge_index.shape[0] == edge_index.shape[1]
81
+ ):
82
+ is_legacy = True
83
+ elif not hasattr(edge_index, "shape"):
84
+ is_legacy = True
85
+
86
+ x_norm = x / (ops.norm(x, axis=-1, keepdims=True) + 1e-12)
87
+
88
+ if is_legacy:
89
+ out = self.propagate(x, edge_index, x_norm=x_norm)
90
+ else:
91
+ num_nodes = x.shape[0] if hasattr(x, "shape") and x.shape[0] is not None else ops.shape(x)[0]
92
+ keep_mask = None
93
+ if self.add_self_loops:
94
+ edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
95
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
96
+ keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
97
+ out = self.propagate(edge_index, x=x, x_norm=x_norm, keep_mask=keep_mask, size=(num_nodes, num_nodes))
98
+
99
+ if self.activation is not None:
100
+ out = self.activation(out)
101
+
102
+ return out
103
+
104
+ def message(self, x=None, x_j=None, x_norm=None, x_norm_i=None, x_norm_j=None, index=None, size_i=None, keep_mask=None):
105
+ beta = self.beta if self.beta is not None else 1.0
106
+
107
+ # Legacy Spektral path
108
+ if x_j is None and x is not None:
109
+ x_j = self.get_sources(x)
110
+ x_norm_i = self.get_targets(x_norm)
111
+ x_norm_j = self.get_sources(x_norm)
112
+ alpha = beta * ops.sum(x_norm_i * x_norm_j, axis=-1)
113
+ alpha = softmax(alpha, self.index_targets, num_nodes=self.n_nodes, dim=0)
114
+ return ops.expand_dims(alpha, -1) * x_j
115
+
116
+ # PyG path
117
+ alpha = beta * ops.sum(x_norm_i * x_norm_j, axis=-1)
118
+ alpha = mask_edge_logits(alpha, keep_mask)
119
+ alpha = softmax(alpha, index, num_nodes=size_i, dim=0)
120
+ return x_j * ops.expand_dims(alpha, -1)
@@ -0,0 +1,94 @@
1
+ from typing import Optional, Union, Callable
2
+ from keras import ops, activations
3
+ from keras.layers import Layer
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.conv.gcn_conv import GCNConv
7
+
8
+
9
+ class AntiSymmetricConv(Layer):
10
+ r"""The anti-symmetric graph convolutional operator from the
11
+ `"Anti-Symmetric DGN: a Continuous approach to Deep Graph Neural Networks"
12
+ <https://arxiv.org/abs/2202.13085>`_ paper.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import AntiSymmetricConv
18
+
19
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
20
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
21
+
22
+ layer = AntiSymmetricConv(in_channels=8)
23
+ out = layer(x, edge_index)
24
+ print(tuple(out.shape)) # (10, 8)
25
+ ```
26
+ """
27
+ def __init__(
28
+ self,
29
+ in_channels: int,
30
+ phi: Optional[MessagePassing] = None,
31
+ num_iters: int = 1,
32
+ epsilon: float = 0.1,
33
+ gamma: float = 0.1,
34
+ act: Union[str, Callable, None] = "tanh",
35
+ bias: bool = True,
36
+ **kwargs,
37
+ ):
38
+ super().__init__(**kwargs)
39
+
40
+ self.in_channels = in_channels
41
+ self.num_iters = num_iters
42
+ self.gamma = gamma
43
+ self.epsilon = epsilon
44
+ self.act = activations.get(act) if act is not None else None
45
+
46
+ if phi is None:
47
+ phi = GCNConv(in_channels, in_channels, bias=False)
48
+ self.phi = phi
49
+
50
+ self.W = self.add_weight(
51
+ shape=(in_channels, in_channels),
52
+ initializer="glorot_uniform",
53
+ name="W",
54
+ )
55
+
56
+ if bias:
57
+ self.bias = self.add_weight(
58
+ shape=(in_channels,),
59
+ initializer="zeros",
60
+ name="bias",
61
+ )
62
+ else:
63
+ self.bias = None
64
+
65
+ def build(self, input_shape=None):
66
+ if hasattr(self.phi, "build"):
67
+ self.phi.build((None, self.in_channels))
68
+ self.built = True
69
+
70
+ def call(self, inputs, edge_index=None, **kwargs):
71
+ if edge_index is None:
72
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
73
+ x, edge_index = inputs
74
+ else:
75
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
76
+ else:
77
+ x = inputs
78
+
79
+ eye = ops.eye(self.in_channels, dtype=self.W.dtype)
80
+ antisymmetric_W = self.W - ops.transpose(self.W) - self.gamma * eye
81
+
82
+ for _ in range(self.num_iters):
83
+ h = self.phi(x, edge_index)
84
+ h = ops.matmul(x, ops.transpose(antisymmetric_W)) + h
85
+
86
+ if self.bias is not None:
87
+ h = h + self.bias
88
+
89
+ if self.act is not None:
90
+ h = self.act(h)
91
+
92
+ x = x + self.epsilon * h
93
+
94
+ return x
@@ -0,0 +1,105 @@
1
+ from keras import layers, ops
2
+
3
+ from k3_node.layers.conv.message_passing import MessagePassing
4
+ from k3_node.layers.conv.utils import gcn_norm, is_tracing
5
+
6
+
7
+ class APPNP(MessagePassing):
8
+ r"""The approximate personalized propagation of neural predictions (APPNP)
9
+ operator from the `"Predict then Propagate: Combining Neural Networks with
10
+ Personalized PageRank for Classification on Graphs"
11
+ <https://arxiv.org/abs/1810.05997>`_ paper.
12
+
13
+ Args:
14
+ K: Number of iterations :math:`K`.
15
+ alpha: Teleport probability :math:`\alpha`.
16
+ dropout: Dropout probability of edges or features during propagation.
17
+ (default: ``0.0``)
18
+ cached: If set to :obj:`True`, the layer will cache the computation of
19
+ normalization coefficients. (default: ``False``)
20
+ add_self_loops: If set to :obj:`False`, will not add self-loops.
21
+ (default: ``True``)
22
+ normalize: Whether to apply symmetric normalization. (default: ``True``)
23
+
24
+ Example:
25
+ ```python
26
+ import numpy as np
27
+ from k3_node.layers import APPNP
28
+
29
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
30
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
31
+
32
+ layer = APPNP(K=2, alpha=0.1)
33
+ out = layer(x, edge_index)
34
+ print(tuple(out.shape)) # (10, 8)
35
+ ```
36
+ """
37
+
38
+ weighted_sum_message = True
39
+
40
+ def __init__(
41
+ self,
42
+ K: int,
43
+ alpha: float,
44
+ dropout: float = 0.0,
45
+ cached: bool = False,
46
+ add_self_loops: bool = True,
47
+ normalize: bool = True,
48
+ **kwargs,
49
+ ):
50
+ super().__init__(aggr="add", **kwargs)
51
+ self.K = K
52
+ self.alpha = alpha
53
+ self.dropout_rate = dropout
54
+ self.cached = cached
55
+ self.add_self_loops = add_self_loops
56
+ self.normalize = normalize
57
+ self._cached_edge_index = None
58
+ self._cached_norm = None
59
+ self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
60
+
61
+ def build(self, input_shape=None):
62
+ if self.dropout is not None and hasattr(self.dropout, "build"):
63
+ self.dropout.build(input_shape)
64
+ self.built = True
65
+
66
+ def call(self, x, edge_index=None, edge_weight=None, training=None, **kwargs):
67
+ if edge_index is None and isinstance(x, (tuple, list)):
68
+ x, edge_index = x[0], x[1]
69
+
70
+ if self.normalize:
71
+ if self.cached and self._cached_edge_index is not None:
72
+ edge_index = self._cached_edge_index
73
+ edge_weight = self._cached_norm
74
+ else:
75
+ num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
76
+ edge_index, edge_weight = gcn_norm(
77
+ edge_index,
78
+ edge_weight,
79
+ num_nodes=num_nodes,
80
+ add_self_loops=self.add_self_loops,
81
+ flow=self.flow,
82
+ dtype=x.dtype,
83
+ )
84
+ if self.cached and not is_tracing(edge_index):
85
+ self._cached_edge_index = edge_index
86
+ self._cached_norm = edge_weight
87
+
88
+ h = x
89
+ for _ in range(self.K):
90
+ if self.dropout is not None:
91
+ h = self.dropout(h, training=training)
92
+ h = self.propagate(edge_index, x=h, edge_weight=edge_weight)
93
+ h = (1.0 - self.alpha) * h + self.alpha * x
94
+
95
+ return h
96
+
97
+ def message(self, x_j, edge_weight=None):
98
+ if edge_weight is None:
99
+ return x_j
100
+ return ops.expand_dims(edge_weight, -1) * x_j
101
+
102
+
103
+ # Backward-compatible alias
104
+ APPNPConv = APPNP
105
+
@@ -0,0 +1,157 @@
1
+ # ported from spektral
2
+
3
+ from keras import activations
4
+ from keras import activations, ops
5
+ from keras.layers import Dense, Dropout
6
+ from keras.models import Sequential
7
+
8
+ from k3_node.layers.conv.conv import Conv
9
+ from k3_node.ops import gcn_filter, modal_dot
10
+
11
+
12
+ class APPNPConv(Conv):
13
+ """
14
+ `k3_node.layers.APPNPConv`
15
+ Implementation of Approximate Personalized Propagation of Neural Predictions
16
+
17
+ Args:
18
+ channels: The number of output channels.
19
+ alpha: The teleport probability.
20
+ propagations: The number of propagation steps.
21
+ mlp_hidden: A list of hidden channels for the MLP.
22
+ mlp_activation: The activation function to use in the MLP.
23
+ dropout_rate: The dropout rate for the MLP.
24
+ activation: The activation function to use in the layer.
25
+ use_bias: Whether to add a bias to the linear transformation.
26
+ kernel_initializer: Initializer for the `kernel` weights matrix.
27
+ bias_initializer: Initializer for the bias vector.
28
+ kernel_regularizer: Regularizer for the `kernel` weights matrix.
29
+ bias_regularizer: Regularizer for the bias vector.
30
+ activity_regularizer: Regularizer for the output.
31
+ kernel_constraint: Constraint for the `kernel` weights matrix.
32
+ bias_constraint: Constraint for the bias vector.
33
+ **kwargs: Additional keyword arguments.
34
+
35
+ Example:
36
+ ```python
37
+ import numpy as np
38
+ from k3_node.layers import APPNPConv
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
+ layer = APPNPConv(channels=16, alpha=0.1, propagations=2)
44
+ out = layer(x, edge_index)
45
+ print(tuple(out.shape)) # (10, 16)
46
+ ```
47
+ """
48
+ def __init__(
49
+ self,
50
+ channels,
51
+ alpha=0.2,
52
+ propagations=1,
53
+ mlp_hidden=None,
54
+ mlp_activation="relu",
55
+ dropout_rate=0.0,
56
+ activation=None,
57
+ use_bias=True,
58
+ kernel_initializer="glorot_uniform",
59
+ bias_initializer="zeros",
60
+ kernel_regularizer=None,
61
+ bias_regularizer=None,
62
+ activity_regularizer=None,
63
+ kernel_constraint=None,
64
+ bias_constraint=None,
65
+ **kwargs,
66
+ ):
67
+
68
+ super().__init__(
69
+ activation=activation,
70
+ use_bias=use_bias,
71
+ kernel_initializer=kernel_initializer,
72
+ bias_initializer=bias_initializer,
73
+ kernel_regularizer=kernel_regularizer,
74
+ bias_regularizer=bias_regularizer,
75
+ activity_regularizer=activity_regularizer,
76
+ kernel_constraint=kernel_constraint,
77
+ bias_constraint=bias_constraint,
78
+ **kwargs,
79
+ )
80
+ self.channels = channels
81
+ self.mlp_hidden = mlp_hidden if mlp_hidden else []
82
+ self.alpha = alpha
83
+ self.propagations = propagations
84
+ self.mlp_activation = activations.get(mlp_activation)
85
+ self.dropout_rate = dropout_rate
86
+
87
+ def build(self, input_shape=None):
88
+ layer_kwargs = dict(
89
+ kernel_initializer=self.kernel_initializer,
90
+ bias_initializer=self.bias_initializer,
91
+ kernel_regularizer=self.kernel_regularizer,
92
+ bias_regularizer=self.bias_regularizer,
93
+ kernel_constraint=self.kernel_constraint,
94
+ bias_constraint=self.bias_constraint,
95
+ dtype=self.dtype,
96
+ )
97
+ mlp_layers = []
98
+ for channels in self.mlp_hidden:
99
+ mlp_layers.extend(
100
+ [
101
+ Dropout(self.dropout_rate),
102
+ Dense(channels, self.mlp_activation, **layer_kwargs),
103
+ ]
104
+ )
105
+ mlp_layers.append(Dense(self.channels, "linear", **layer_kwargs))
106
+ self.mlp = Sequential(mlp_layers)
107
+ if input_shape is not None:
108
+ feat_shape = input_shape[0] if isinstance(input_shape, (list, tuple)) else input_shape
109
+ self.mlp.build(feat_shape)
110
+ self.built = True
111
+
112
+ def call(self, inputs, mask=None):
113
+ x, a = inputs
114
+ def call(self, inputs, a=None, mask=None):
115
+ if a is not None:
116
+ x = inputs
117
+ elif isinstance(inputs, (list, tuple)) and len(inputs) == 2:
118
+ x, a = inputs
119
+ else:
120
+ x = inputs
121
+ a = None
122
+
123
+ if not self.built:
124
+ self.build(getattr(x, "shape", None))
125
+
126
+ if a is not None and hasattr(a, "shape") and len(a.shape) == 2 and a.shape[0] == 2 and a.shape[1] != 2:
127
+ num_nodes = ops.shape(x)[-2] if len(ops.shape(x)) >= 2 else ops.shape(x)[0]
128
+ row, col = ops.cast(a[0], "int32"), ops.cast(a[1], "int32")
129
+ idx = ops.stack([row, col], axis=-1)
130
+ zeros = ops.zeros((num_nodes, num_nodes), dtype=x.dtype)
131
+ a = ops.scatter_update(zeros, idx, ops.ones((ops.shape(a)[1],), dtype=x.dtype))
132
+
133
+ mlp_out = self.mlp(x)
134
+ output = mlp_out
135
+ if a is not None:
136
+ for _ in range(self.propagations):
137
+ output = (1 - self.alpha) * modal_dot(a, output) + self.alpha * mlp_out
138
+ if mask is not None and isinstance(mask, (list, tuple)) and len(mask) > 0 and mask[0] is not None:
139
+ output *= mask[0]
140
+ output = self.activation(output)
141
+
142
+ return output
143
+
144
+ @property
145
+ def config(self):
146
+ return {
147
+ "channels": self.channels,
148
+ "alpha": self.alpha,
149
+ "propagations": self.propagations,
150
+ "mlp_hidden": self.mlp_hidden,
151
+ "mlp_activation": activations.serialize(self.mlp_activation),
152
+ "dropout_rate": self.dropout_rate,
153
+ }
154
+
155
+ @staticmethod
156
+ def preprocess(a):
157
+ return gcn_filter(a)