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,144 @@
1
+ # ported from stellargraph
2
+ from keras import ops
3
+ from keras import activations, initializers, constraints, regularizers
4
+ from keras.layers import Layer, dot
5
+ from k3_node.ops.creation import repeat
6
+
7
+
8
+ class GraphConvolution(Layer):
9
+ """
10
+ `k3_node.layers.GraphConvolution`
11
+ Implementation of Graph Convolution (GCN) layer
12
+
13
+ Args:
14
+ units: Positive integer, dimensionality of the output space.
15
+ activation: Activation function to use.
16
+ use_bias: Whether to add a bias to the linear transformation.
17
+ final_layer: Deprecated, use tf.gather or GatherIndices instead.
18
+ input_dim: Deprecated, use `keras.layers.Input` with `input_shape` instead.
19
+ kernel_initializer: Initializer for the `kernel` weights matrix.
20
+ kernel_regularizer: Regularizer for the `kernel` weights matrix.
21
+ kernel_constraint: Constraint for the `kernel` weights matrix.
22
+ bias_initializer: Initializer for the bias vector.
23
+ bias_regularizer: Regularizer for the bias vector.
24
+ bias_constraint: Constraint for the bias vector.
25
+ **kwargs: Additional arguments to pass to the `Layer` superclass.
26
+
27
+ Example:
28
+ ```python
29
+ import numpy as np
30
+ from k3_node.layers import GraphConvolution
31
+
32
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
33
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
34
+
35
+ layer = GraphConvolution(units=16, activation="relu")
36
+ out = layer(x, edge_index)
37
+ print(tuple(out.shape)) # (10, 16)
38
+ ```
39
+ """
40
+ def __init__(
41
+ self,
42
+ units,
43
+ activation=None,
44
+ use_bias=True,
45
+ final_layer=None,
46
+ input_dim=None,
47
+ kernel_initializer="glorot_uniform",
48
+ kernel_regularizer=None,
49
+ kernel_constraint=None,
50
+ bias_initializer="zeros",
51
+ bias_regularizer=None,
52
+ bias_constraint=None,
53
+ **kwargs,
54
+ ):
55
+ if isinstance(activation, int):
56
+ # Called as GraphConvolution(in_channels, out_channels)
57
+ self.in_channels = units
58
+ units = activation
59
+ activation = kwargs.pop("activation", None)
60
+
61
+ if "input_shape" not in kwargs and input_dim is not None:
62
+ kwargs["input_shape"] = (input_dim,)
63
+
64
+ self.units = units
65
+ self.activation = activations.get(activation)
66
+ self.use_bias = use_bias
67
+ if final_layer is not None:
68
+ raise ValueError(
69
+ "'final_layer' is not longer supported, use 'tf.gather' or 'GatherIndices' separately"
70
+ )
71
+
72
+ self.kernel_initializer = initializers.get(kernel_initializer)
73
+ self.kernel_regularizer = regularizers.get(kernel_regularizer)
74
+ self.kernel_constraint = constraints.get(kernel_constraint)
75
+ self.bias_initializer = initializers.get(bias_initializer)
76
+ self.bias_regularizer = regularizers.get(bias_regularizer)
77
+ self.bias_constraint = constraints.get(bias_constraint)
78
+
79
+ super().__init__(**kwargs)
80
+
81
+ def build(self, input_shapes):
82
+ if isinstance(input_shapes, (list, tuple)) and len(input_shapes) > 0 and isinstance(input_shapes[0], (list, tuple)):
83
+ feat_shape = input_shapes[0]
84
+ else:
85
+ feat_shape = input_shapes
86
+ input_dim = int(feat_shape[-1]) if feat_shape is not None and feat_shape[-1] is not None else 8
87
+
88
+ self.kernel = self.add_weight(
89
+ shape=(1, input_dim, self.units),
90
+ initializer=self.kernel_initializer,
91
+ name="kernel",
92
+ regularizer=self.kernel_regularizer,
93
+ constraint=self.kernel_constraint,
94
+ )
95
+
96
+ if self.use_bias:
97
+ self.bias = self.add_weight(
98
+ shape=(self.units,),
99
+ initializer=self.bias_initializer,
100
+ name="bias",
101
+ regularizer=self.bias_regularizer,
102
+ constraint=self.bias_constraint,
103
+ )
104
+ else:
105
+ self.bias = None
106
+ self.built = True
107
+
108
+ def call(self, inputs, A=None, **kwargs):
109
+ if A is not None:
110
+ features = inputs
111
+ elif isinstance(inputs, (list, tuple)):
112
+ features, A = inputs
113
+ else:
114
+ features, A = inputs, None
115
+
116
+ if A is not None and hasattr(A, "shape") and len(A.shape) == 2 and A.shape[0] == 2 and A.shape[1] != 2:
117
+ num_nodes = ops.shape(features)[-2]
118
+ a_dense = ops.zeros((num_nodes, num_nodes), dtype=features.dtype)
119
+ indices = ops.transpose(A, axes=[1, 0])
120
+ updates = ops.ones(shape=(ops.shape(A)[1],), dtype=features.dtype)
121
+ A = ops.scatter_update(a_dense, indices, updates)
122
+
123
+ was_2d = len(ops.shape(features)) == 2
124
+ if was_2d:
125
+ features = ops.expand_dims(features, 0)
126
+ if len(ops.shape(A)) == 2:
127
+ A = ops.expand_dims(A, 0)
128
+
129
+ # Calculate the layer operation of GCN
130
+
131
+ h_graph = dot((A, features), axes=1)
132
+ b = ops.shape(h_graph)[0]
133
+ kernel = repeat(self.kernel, b, axis=0)
134
+ output = dot((h_graph, kernel), axes=(-1, 1))
135
+
136
+ # Add optional bias & apply activation
137
+ if self.bias is not None:
138
+ output += self.bias
139
+ output = self.activation(output)
140
+
141
+ if was_2d:
142
+ output = ops.squeeze(output, 0)
143
+
144
+ return output
@@ -0,0 +1,126 @@
1
+ import math
2
+ from typing import Optional, Union, Tuple
3
+ from keras import ops
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.conv.utils import gcn_norm, is_tracing
7
+
8
+
9
+ class GCN2Conv(MessagePassing):
10
+ r"""The graph convolutional operator from the `"Simple and Deep Graph
11
+ Convolutional Networks" <https://arxiv.org/abs/2007.02133>`_ paper.
12
+
13
+ Args:
14
+ channels: Size of each input and output sample.
15
+ alpha: The strength of the initial residual connection :math:`\alpha`.
16
+ theta: The hyperparameter for the identity mapping :math:`\theta`.
17
+ (default: :obj:`None`)
18
+ layer: The layer index :math:`l`. (default: :obj:`None`)
19
+ shared_weights: If set to :obj:`True`, will use the same weights
20
+ for :math:`\mathbf{X}` and :math:`\mathbf{X}_0`. (default: :obj:`True`)
21
+ cached: If set to :obj:`True`, will cache the computation of normalization
22
+ coefficients. (default: ``False``)
23
+ add_self_loops: If set to :obj:`False`, will not add self-loops.
24
+ (default: ``True``)
25
+ normalize: Whether to apply symmetric normalization. (default: ``True``)
26
+
27
+ Example:
28
+ ```python
29
+ import numpy as np
30
+ from k3_node.layers import GCN2Conv
31
+
32
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
33
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
34
+
35
+ x_0 = x # initial node representations (from the first layer)
36
+ layer = GCN2Conv(channels=8, alpha=0.1, theta=0.5, layer=1)
37
+ out = layer(x, x_0, edge_index)
38
+ print(tuple(out.shape)) # (10, 8)
39
+ ```
40
+ """
41
+
42
+ weighted_sum_message = True
43
+
44
+ def __init__(
45
+ self,
46
+ channels: int,
47
+ alpha: float,
48
+ theta: Optional[float] = None,
49
+ layer: Optional[int] = None,
50
+ shared_weights: bool = True,
51
+ cached: bool = False,
52
+ add_self_loops: bool = True,
53
+ normalize: bool = True,
54
+ **kwargs,
55
+ ):
56
+ super().__init__(aggr="add", **kwargs)
57
+ self.channels = channels
58
+ self.alpha = alpha
59
+ self.beta = 1.0
60
+ if theta is not None and layer is not None:
61
+ self.beta = math.log(theta / layer + 1.0)
62
+ self.cached = cached
63
+ self.normalize = normalize
64
+ self.add_self_loops = add_self_loops
65
+ self.shared_weights = shared_weights
66
+
67
+ self._cached_edge_index = None
68
+ self._cached_norm = None
69
+
70
+ def build(self, input_shape):
71
+ self.weight1 = self.add_weight(
72
+ shape=(self.channels, self.channels),
73
+ initializer="glorot_uniform",
74
+ name="weight1",
75
+ )
76
+ if not self.shared_weights:
77
+ self.weight2 = self.add_weight(
78
+ shape=(self.channels, self.channels),
79
+ initializer="glorot_uniform",
80
+ name="weight2",
81
+ )
82
+ else:
83
+ self.weight2 = None
84
+ self.built = True
85
+
86
+ def call(self, x, x_0, edge_index=None, edge_weight=None, **kwargs):
87
+ if edge_index is None and isinstance(x, (tuple, list)):
88
+ x, x_0, edge_index = x[0], x[1], x[2]
89
+
90
+ if self.normalize:
91
+ if self.cached and self._cached_edge_index is not None:
92
+ edge_index = self._cached_edge_index
93
+ edge_weight = self._cached_norm
94
+ else:
95
+ 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]
96
+ edge_index, edge_weight = gcn_norm(
97
+ edge_index,
98
+ edge_weight,
99
+ num_nodes=num_nodes,
100
+ add_self_loops=self.add_self_loops,
101
+ flow=self.flow,
102
+ dtype=x.dtype,
103
+ )
104
+ if self.cached and not is_tracing(edge_index):
105
+ self._cached_edge_index = edge_index
106
+ self._cached_norm = edge_weight
107
+
108
+ h = self.propagate(edge_index, x=x, edge_weight=edge_weight)
109
+ h = (1.0 - self.alpha) * h
110
+ h_0 = self.alpha * x_0
111
+
112
+ if self.weight2 is None:
113
+ combined = h + h_0
114
+ out = (1.0 - self.beta) * combined + self.beta * ops.matmul(combined, self.weight1)
115
+ else:
116
+ term1 = (1.0 - self.beta) * h + self.beta * ops.matmul(h, self.weight1)
117
+ term2 = (1.0 - self.beta) * h_0 + self.beta * ops.matmul(h_0, self.weight2)
118
+ out = term1 + term2
119
+
120
+ return out
121
+
122
+ def message(self, x_j, edge_weight=None):
123
+ if edge_weight is None:
124
+ return x_j
125
+ return ops.expand_dims(edge_weight, -1) * x_j
126
+
@@ -0,0 +1,135 @@
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 GCNConv(MessagePassing):
8
+ r"""The graph convolutional operator from the `"Semi-supervised
9
+ Classification with Graph Convolutional Networks"
10
+ <https://arxiv.org/abs/1609.02907>`_ paper.
11
+
12
+ .. math::
13
+ \mathbf{X}^{\prime} = \mathbf{\hat{D}}^{-1/2} \mathbf{\hat{A}}
14
+ \mathbf{\hat{D}}^{-1/2} \mathbf{X} \mathbf{\Theta}
15
+
16
+ Args:
17
+ in_channels: Size of each input sample.
18
+ out_channels: Size of each output sample.
19
+ improved: If set to :obj:`True`, the layer computes
20
+ :math:`\mathbf{\hat{A}} = \mathbf{A} + 2 \mathbf{I}`.
21
+ (default: :obj:`False`)
22
+ cached: If set to :obj:`True`, the layer will cache the computation of
23
+ :math:`\mathbf{\hat{D}}^{-1/2} \mathbf{\hat{A}} \mathbf{\hat{D}}^{-1/2}`.
24
+ (default: :obj:`False`)
25
+ add_self_loops: If set to :obj:`False`, will not add
26
+ self-loops to the input graph. (default: :obj:`True`)
27
+ normalize: Whether to add self-loops and compute
28
+ symmetric normalization coefficients on the fly.
29
+ (default: :obj:`True`)
30
+ bias: If set to :obj:`False`, the layer will not learn
31
+ an additive bias. (default: :obj:`True`)
32
+
33
+ Example:
34
+ ```python
35
+ import numpy as np
36
+ from k3_node.layers import GCNConv
37
+
38
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
39
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
40
+
41
+ layer = GCNConv(in_channels=8, out_channels=16)
42
+ out = layer(x, edge_index)
43
+ print(tuple(out.shape)) # (10, 16)
44
+ ```
45
+ """
46
+
47
+ weighted_sum_message = True
48
+
49
+ def __init__(
50
+ self,
51
+ in_channels: int,
52
+ out_channels: int,
53
+ improved: bool = False,
54
+ cached: bool = False,
55
+ add_self_loops: bool = True,
56
+ normalize: bool = True,
57
+ bias: bool = True,
58
+ **kwargs,
59
+ ):
60
+ super().__init__(aggr="add", **kwargs)
61
+ self.in_channels = in_channels
62
+ self.out_channels = out_channels
63
+ self.improved = improved
64
+ self.cached = cached
65
+ self.add_self_loops = add_self_loops
66
+ self.normalize = normalize
67
+ self.use_bias = bias
68
+
69
+ self.lin = layers.Dense(out_channels, use_bias=False)
70
+ self.bias = None
71
+ self._cached_edge_index = None
72
+ self._cached_norm = None
73
+
74
+ def build(self, input_shape):
75
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
76
+ feat_shape = input_shape[0]
77
+ else:
78
+ feat_shape = input_shape
79
+ self.lin.build(feat_shape)
80
+ if self.use_bias:
81
+ self.bias = self.add_weight(
82
+ shape=(self.out_channels,),
83
+ initializer="zeros",
84
+ name="bias",
85
+ )
86
+ else:
87
+ self.bias = None
88
+ self.built = True
89
+
90
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
91
+ # Handle legacy call conv((x, adj))
92
+ if edge_index is None and isinstance(x, (tuple, list)):
93
+ x, edge_index = x[0], x[1]
94
+
95
+ if not self.built:
96
+ feat_shape = x.shape if hasattr(x, "shape") and x.shape is not None else (None, self.in_channels)
97
+ self.build(feat_shape)
98
+
99
+ # Handle dense adjacency [N, N]
100
+ e_shape = getattr(edge_index, "shape", None)
101
+ if e_shape is not None and len(e_shape) == 2 and e_shape[0] is not None and e_shape[1] is not None and e_shape[0] > 2 and e_shape[0] == e_shape[1]:
102
+ where_adj = ops.where(edge_index != 0)
103
+ where_adj = where_adj if not isinstance(where_adj, list) else where_adj
104
+ edge_weight = ops.take(edge_index, where_adj[0] * e_shape[1] + where_adj[1]) if edge_weight is None else edge_weight
105
+ edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
106
+
107
+ if self.normalize:
108
+ if self.cached and self._cached_edge_index is not None:
109
+ edge_index = self._cached_edge_index
110
+ edge_weight = self._cached_norm
111
+ else:
112
+ 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]
113
+ edge_index, edge_weight = gcn_norm(
114
+ edge_index,
115
+ edge_weight,
116
+ num_nodes=num_nodes,
117
+ improved=self.improved,
118
+ add_self_loops=self.add_self_loops,
119
+ flow=self.flow,
120
+ dtype=x.dtype,
121
+ )
122
+ if self.cached and not is_tracing(edge_index):
123
+ self._cached_edge_index = edge_index
124
+ self._cached_norm = edge_weight
125
+
126
+ x = self.lin(x)
127
+ out = self.propagate(edge_index, x=x, edge_weight=edge_weight)
128
+ if self.bias is not None:
129
+ out = out + self.bias
130
+ return out
131
+
132
+ def message(self, x_j, edge_weight=None):
133
+ if edge_weight is None:
134
+ return x_j
135
+ return ops.expand_dims(edge_weight, -1) * x_j
@@ -0,0 +1,163 @@
1
+ from typing import Optional, Union, Tuple
2
+ import keras
3
+ from keras import ops
4
+ from keras.layers import Dense, BatchNormalization, LayerNormalization
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+ from k3_node.layers.aggr import SoftmaxAggregation, PowerMeanAggregation
8
+
9
+
10
+ class GENConv(MessagePassing):
11
+ r"""The generalized graph convolution operator from the `"DeeperGCN: All
12
+ You Need to Train Deeper GCNs" <https://arxiv.org/abs/2006.07739>`_ paper.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import GENConv
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
+ edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
22
+
23
+ layer = GENConv(in_channels=8, out_channels=16, edge_dim=3)
24
+ out = layer(x, edge_index, edge_attr)
25
+ print(tuple(out.shape)) # (10, 16)
26
+ ```
27
+ """
28
+ def __init__(
29
+ self,
30
+ in_channels: Union[int, Tuple[int, int]],
31
+ out_channels: int,
32
+ aggr: str = "softmax",
33
+ t: float = 1.0,
34
+ learn_t: bool = False,
35
+ p: float = 1.0,
36
+ learn_p: bool = False,
37
+ msg_norm: bool = False,
38
+ learn_msg_scale: bool = False,
39
+ norm: Optional[str] = "batch",
40
+ num_layers: int = 2,
41
+ expansion: int = 2,
42
+ eps: float = 1e-7,
43
+ bias: bool = False,
44
+ edge_dim: Optional[int] = None,
45
+ **kwargs,
46
+ ):
47
+ if aggr in ("softmax", "softmax_sg"):
48
+ aggr_module = SoftmaxAggregation(t=t, learn=learn_t)
49
+ elif aggr in ("power", "powermean"):
50
+ aggr_module = PowerMeanAggregation(p=p, learn=learn_p)
51
+ else:
52
+ aggr_module = aggr
53
+
54
+ super().__init__(aggr=aggr_module, **kwargs)
55
+
56
+ self.in_channels = in_channels
57
+ self.out_channels = out_channels
58
+ self.eps = eps
59
+ self.edge_dim = edge_dim
60
+ self.use_bias = bias
61
+
62
+ if isinstance(in_channels, int):
63
+ self.in_channels_l = in_channels
64
+ self.in_channels_r = in_channels
65
+ else:
66
+ self.in_channels_l, self.in_channels_r = in_channels
67
+
68
+ if self.in_channels_l != out_channels:
69
+ self.lin_src = Dense(out_channels, use_bias=bias)
70
+ else:
71
+ self.lin_src = None
72
+
73
+ if edge_dim is not None and edge_dim != out_channels:
74
+ self.lin_edge = Dense(out_channels, use_bias=bias)
75
+ else:
76
+ self.lin_edge = None
77
+
78
+ if self.in_channels_r != out_channels:
79
+ self.lin_dst = Dense(out_channels, use_bias=bias)
80
+ else:
81
+ self.lin_dst = None
82
+
83
+ # MLP
84
+ self.mlp_layers = []
85
+ channels = [out_channels]
86
+ for _ in range(num_layers - 1):
87
+ channels.append(out_channels * expansion)
88
+ channels.append(out_channels)
89
+
90
+ for i in range(len(channels) - 1):
91
+ self.mlp_layers.append(Dense(channels[i + 1], use_bias=bias))
92
+ if i < len(channels) - 2:
93
+ if norm == "batch":
94
+ self.mlp_layers.append(BatchNormalization(momentum=0.9, epsilon=1e-5))
95
+ elif norm == "layer":
96
+ self.mlp_layers.append(LayerNormalization())
97
+ self.mlp_layers.append(keras.layers.ReLU())
98
+
99
+ def build(self, input_shape=None):
100
+ if self.lin_src is not None:
101
+ self.lin_src.build((None, self.in_channels_l))
102
+ if self.lin_dst is not None:
103
+ self.lin_dst.build((None, self.in_channels_r))
104
+ if self.lin_edge is not None:
105
+ self.lin_edge.build((None, self.edge_dim))
106
+ curr_dim = self.out_channels
107
+ for layer in self.mlp_layers:
108
+ if hasattr(layer, "build"):
109
+ layer.build((None, curr_dim))
110
+ if hasattr(layer, "units"):
111
+ curr_dim = layer.units
112
+ self.built = True
113
+
114
+ def call(self, inputs, edge_index=None, edge_attr=None, training=None, **kwargs):
115
+ if edge_index is None:
116
+ if isinstance(inputs, (list, tuple)):
117
+ if len(inputs) == 3:
118
+ x, edge_index, edge_attr = inputs
119
+ elif len(inputs) == 2:
120
+ x, edge_index = inputs
121
+ else:
122
+ raise ValueError(f"Unexpected input length {len(inputs)}")
123
+ else:
124
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
125
+ else:
126
+ x = inputs
127
+
128
+ if not self.built:
129
+ self.build()
130
+
131
+ if isinstance(x, (list, tuple)):
132
+ x_l, x_r = x
133
+ else:
134
+ x_l = x_r = x
135
+
136
+ if self.lin_src is not None:
137
+ x_l = self.lin_src(x_l)
138
+
139
+ num_nodes = ops.shape(x_r)[0]
140
+ out = self.propagate(
141
+ edge_index,
142
+ x=(x_l, x_r),
143
+ edge_attr=edge_attr,
144
+ size=(ops.shape(x_l)[0], num_nodes),
145
+ )
146
+
147
+ x_dst = x_r
148
+ if self.lin_dst is not None:
149
+ x_dst = self.lin_dst(x_dst)
150
+ out = out + x_dst
151
+
152
+ for layer in self.mlp_layers:
153
+ # Batch norm needs `training` explicitly (Keras does not propagate it on JAX).
154
+ out = layer(out, training=training) if isinstance(layer, BatchNormalization) else layer(out)
155
+
156
+ return out
157
+
158
+ def message(self, x_j, edge_attr=None):
159
+ if edge_attr is not None and self.lin_edge is not None:
160
+ edge_attr = self.lin_edge(edge_attr)
161
+ msg = x_j if edge_attr is None else x_j + edge_attr
162
+ return ops.relu(msg) + self.eps
163
+