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,218 @@
1
+ import math
2
+ from typing import Optional, Union, Tuple
3
+ from keras import ops
4
+ from keras.layers import Dense
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+ from k3_node.layers.conv.utils import softmax
8
+
9
+
10
+ class GeneralConv(MessagePassing):
11
+ r"""A general GNN layer adapted from the `"Design Space for Graph Neural
12
+ Networks" <https://arxiv.org/abs/2011.08843>`_ paper.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import GeneralConv
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 = GeneralConv(in_channels=8, out_channels=16, in_edge_channels=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: Optional[int] = None,
32
+ in_edge_channels: Optional[int] = None,
33
+ aggr: str = "add",
34
+ skip_linear: bool = False,
35
+ directed_msg: bool = True,
36
+ heads: int = 1,
37
+ attention: bool = False,
38
+ attention_type: str = "additive",
39
+ l2_normalize: bool = False,
40
+ bias: bool = True,
41
+ # Spektral compatibility arguments:
42
+ channels: Optional[int] = None,
43
+ batch_norm: Optional[bool] = None,
44
+ dropout: Optional[float] = None,
45
+ aggregate: Optional[str] = None,
46
+ activation: Optional[str] = None,
47
+ use_bias: Optional[bool] = None,
48
+ **kwargs,
49
+ ):
50
+ if out_channels is None:
51
+ if channels is not None:
52
+ out_channels = channels
53
+ in_channels = -1
54
+ else:
55
+ out_channels = in_channels
56
+ in_channels = -1
57
+
58
+ if aggregate is not None:
59
+ aggr = aggregate
60
+ if use_bias is not None:
61
+ bias = use_bias
62
+
63
+ kwargs.setdefault("aggr", aggr)
64
+ super().__init__(node_dim=0, **kwargs)
65
+
66
+ self.in_channels = in_channels
67
+ self.out_channels = out_channels
68
+ self.in_edge_channels = in_edge_channels
69
+ self.skip_linear = skip_linear
70
+ self.directed_msg = directed_msg
71
+ self.heads = heads
72
+ self.attention = attention
73
+ self.attention_type = attention_type
74
+ self.l2_normalize = l2_normalize
75
+ self.use_bias = bias
76
+
77
+ if isinstance(in_channels, int):
78
+ self.in_channels_l = in_channels
79
+ self.in_channels_r = in_channels
80
+ else:
81
+ self.in_channels_l, self.in_channels_r = in_channels
82
+
83
+ self.lin_msg = Dense(out_channels * heads, use_bias=bias)
84
+ if not directed_msg:
85
+ self.lin_msg_i = Dense(out_channels * heads, use_bias=bias)
86
+ else:
87
+ self.lin_msg_i = None
88
+
89
+ if skip_linear or self.in_channels_r != out_channels:
90
+ self.lin_self = Dense(out_channels, use_bias=bias)
91
+ else:
92
+ self.lin_self = None
93
+
94
+ if in_edge_channels is not None:
95
+ self.lin_edge = Dense(out_channels * heads, use_bias=bias)
96
+ else:
97
+ self.lin_edge = None
98
+
99
+ if attention:
100
+ if attention_type == "additive":
101
+ self.att_msg = self.add_weight(
102
+ shape=(1, heads, out_channels),
103
+ initializer="glorot_uniform",
104
+ name="att_msg",
105
+ )
106
+ elif attention_type == "dot_product":
107
+ self.scaler = math.sqrt(out_channels)
108
+ else:
109
+ raise ValueError(f"Attention type '{attention_type}' not supported")
110
+ else:
111
+ self.att_msg = None
112
+
113
+ def build(self, input_shape=None):
114
+ if input_shape is not None:
115
+ if isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and isinstance(input_shape[0], (list, tuple)):
116
+ dim = input_shape[0][-1]
117
+ elif isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and input_shape[0] is not None and not isinstance(input_shape[0], (int, type(None))):
118
+ dim = getattr(input_shape[0], "shape", [None, None])[-1]
119
+ else:
120
+ dim = input_shape[-1]
121
+ if (self.in_channels_l is None or self.in_channels_l == -1) and dim is not None:
122
+ self.in_channels_l = self.in_channels_r = dim
123
+ if self.in_channels_l is not None and self.in_channels_l != -1:
124
+ self.lin_msg.build((None, self.in_channels_l))
125
+ if self.lin_msg_i is not None:
126
+ self.lin_msg_i.build((None, self.in_channels_r))
127
+ if self.lin_self is not None:
128
+ self.lin_self.build((None, self.in_channels_r))
129
+ if self.lin_edge is not None:
130
+ self.lin_edge.build((None, self.in_edge_channels))
131
+ self.built = True
132
+
133
+ def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
134
+ if edge_index is None:
135
+ if isinstance(inputs, (list, tuple)):
136
+ if len(inputs) == 3:
137
+ x, edge_index, edge_attr = inputs
138
+ elif len(inputs) == 2:
139
+ x, edge_index = inputs
140
+ else:
141
+ raise ValueError(f"Unexpected inputs length {len(inputs)}")
142
+ else:
143
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
144
+ else:
145
+ x = inputs
146
+
147
+ if isinstance(x, (list, tuple)):
148
+ x_l, x_r = x
149
+ else:
150
+ x_l = x_r = x
151
+
152
+ if self.in_channels_l is None or self.in_channels_l == -1:
153
+ dim = ops.shape(x_l)[-1]
154
+ self.in_channels_l = self.in_channels_r = dim
155
+ self.build((None, dim))
156
+
157
+ # Check legacy adj matrix
158
+ is_legacy = False
159
+ if hasattr(edge_index, "shape") and len(edge_index.shape) == 2:
160
+ if edge_index.shape[0] != 2 and edge_index.shape[0] == edge_index.shape[1]:
161
+ is_legacy = True
162
+ elif not hasattr(edge_index, "shape"):
163
+ is_legacy = True
164
+
165
+ if is_legacy:
166
+ if hasattr(edge_index, "indices") and not callable(edge_index.indices):
167
+ edge_index = ops.transpose(edge_index.indices)
168
+ else:
169
+ adj = edge_index
170
+ row, col = ops.where(adj > 0)
171
+ edge_index = ops.stack([row, col], axis=0)
172
+
173
+ num_nodes = ops.shape(x_r)[0]
174
+ size = (ops.shape(x_l)[0], num_nodes)
175
+
176
+ out = self.propagate(
177
+ edge_index,
178
+ x=(x_l, x_r),
179
+ edge_attr=edge_attr,
180
+ size=size,
181
+ )
182
+ out = ops.mean(out, axis=1) # aggregate heads
183
+
184
+ if self.lin_self is not None:
185
+ out = out + self.lin_self(x_r)
186
+ else:
187
+ out = out + x_r
188
+
189
+ if self.l2_normalize:
190
+ out = out / (ops.norm(out, axis=-1, keepdims=True) + 1e-12)
191
+
192
+ return out
193
+
194
+ def _message_basic(self, x_i, x_j, edge_attr=None):
195
+ if self.directed_msg:
196
+ x_j = self.lin_msg(x_j)
197
+ else:
198
+ x_j = self.lin_msg(x_j) + self.lin_msg_i(x_i)
199
+ if edge_attr is not None and self.lin_edge is not None:
200
+ x_j = x_j + self.lin_edge(edge_attr)
201
+ return x_j
202
+
203
+ def message(self, x_i, x_j, edge_attr=None, index=None, size_i=None):
204
+ x_j_out = self._message_basic(x_i, x_j, edge_attr)
205
+ x_j_out = ops.reshape(x_j_out, (-1, self.heads, self.out_channels))
206
+
207
+ if self.attention:
208
+ if self.attention_type == "dot_product":
209
+ x_i_out = self._message_basic(x_j, x_i, edge_attr)
210
+ x_i_out = ops.reshape(x_i_out, (-1, self.heads, self.out_channels))
211
+ alpha = ops.sum(x_i_out * x_j_out, axis=-1) / self.scaler
212
+ else:
213
+ alpha = ops.sum(x_j_out * self.att_msg, axis=-1)
214
+ alpha = ops.leaky_relu(alpha, negative_slope=0.2)
215
+ alpha = softmax(alpha, index, num_nodes=size_i, dim=0)
216
+ return x_j_out * ops.expand_dims(alpha, -1)
217
+ else:
218
+ return x_j_out
@@ -0,0 +1,218 @@
1
+ from typing import Callable, Optional, Union, List
2
+ import keras
3
+ from keras import layers, ops, activations
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+
7
+
8
+ def _apply_nn(nn, x, training):
9
+ # Forward `training` explicitly: on the JAX backend Keras does not propagate it to nested
10
+ # layers, so dropout / batch norm inside `nn` would otherwise never be in training mode.
11
+ if isinstance(nn, keras.layers.Layer):
12
+ return nn(x, training=training)
13
+ return nn(x)
14
+
15
+ class GINConv(MessagePassing):
16
+ r"""The graph isomorphism operator from the `"How Powerful are Graph
17
+ Neural Networks?" <https://arxiv.org/abs/1810.00826>`_ paper.
18
+
19
+ Args:
20
+ nn: A neural network :math:`h_{\mathbf{\Theta}}` that maps node
21
+ features to new embeddings (e.g. a :class:`keras.Sequential` or callable).
22
+ Also accepts an integer `channels` for backward compatibility.
23
+ eps: (Initial) :math:`\epsilon`-value. (default: ``0.0``)
24
+ train_eps: If set to :obj:`True`, :math:`\epsilon` will be a learnable
25
+ parameter. (default: ``False``)
26
+
27
+ Example:
28
+ ```python
29
+ import numpy as np
30
+ import keras
31
+ from k3_node.layers import GINConv
32
+
33
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
34
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
35
+
36
+ mlp = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
37
+ layer = GINConv(mlp)
38
+ out = layer(x, edge_index)
39
+ print(tuple(out.shape)) # (10, 16)
40
+ ```
41
+ """
42
+
43
+ def __init__(
44
+ self,
45
+ nn: Union[Callable, int],
46
+ eps: float = 0.0,
47
+ train_eps: bool = False,
48
+ epsilon: Optional[float] = None,
49
+ mlp_hidden: Optional[List[int]] = None,
50
+ mlp_activation: str = "relu",
51
+ mlp_batchnorm: bool = True,
52
+ **kwargs,
53
+ ):
54
+ super().__init__(aggr=kwargs.pop("aggr", kwargs.pop("aggregate", "add")), **kwargs)
55
+
56
+ if epsilon is not None:
57
+ eps = epsilon
58
+
59
+ self.initial_eps = eps
60
+ self.train_eps = train_eps
61
+
62
+ # Backward compatibility if nn is int (channels)
63
+ if isinstance(nn, int):
64
+ channels = nn
65
+ mlp_hidden = mlp_hidden or []
66
+ act = activations.get(mlp_activation)
67
+ seq_layers = []
68
+ for h in mlp_hidden:
69
+ seq_layers.append(layers.Dense(h, activation=act))
70
+ if mlp_batchnorm:
71
+ seq_layers.append(layers.BatchNormalization(momentum=0.9, epsilon=1e-5))
72
+ seq_layers.append(layers.Dense(channels, activation=kwargs.get("activation", None)))
73
+ self.nn = keras.Sequential(seq_layers)
74
+ else:
75
+ self.nn = nn
76
+
77
+ if train_eps:
78
+ self.eps = self.add_weight(
79
+ shape=(1,),
80
+ initializer=keras.initializers.Constant(eps),
81
+ name="eps",
82
+ )
83
+ else:
84
+ self.eps = ops.cast(eps, "float32")
85
+
86
+ def build(self, input_shape):
87
+ if hasattr(self.nn, "build") and not getattr(self.nn, "built", False):
88
+ self.nn.build(input_shape)
89
+ self.built = True
90
+
91
+ def call(self, x, edge_index=None, size=None, training=None, **kwargs):
92
+ # Handle legacy calling: conv((x, adj))
93
+ if edge_index is None and isinstance(x, (tuple, list)) and len(x) == 2:
94
+ arg0, arg1 = x[0], x[1]
95
+ s1 = getattr(arg1, "shape", None)
96
+ if (
97
+ s1 is not None
98
+ and len(s1) == 2
99
+ and s1[0] is not None
100
+ and s1[1] is not None
101
+ and s1[0] > 2
102
+ and s1[0] == s1[1]
103
+ ):
104
+ where_adj = ops.where(arg1 != 0)
105
+ where_adj = where_adj if not isinstance(where_adj, list) else where_adj
106
+ edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
107
+ x = arg0
108
+ elif s1 is not None and len(s1) >= 1 and s1[0] == 2:
109
+ edge_index = arg1
110
+ x = arg0
111
+
112
+ if not isinstance(x, (tuple, list)):
113
+ x_src, x_dst = x, x
114
+ else:
115
+ x_src, x_dst = x[0], x[1]
116
+
117
+ out = self.propagate(edge_index, x=(x_src, x_dst), size=size)
118
+
119
+ if x_dst is not None:
120
+ out = out + (1.0 + self.eps) * x_dst
121
+
122
+ return _apply_nn(self.nn, out, training)
123
+
124
+ def message(self, x_j):
125
+ return x_j
126
+
127
+
128
+ class GINEConv(MessagePassing):
129
+ r"""The modified :class:`GINConv` operator from the `"Strategies for
130
+ Pre-training Graph Neural Networks" <https://arxiv.org/abs/1905.12265>`_
131
+ paper, which is able to incorporate edge features into aggregation.
132
+
133
+ Args:
134
+ nn: A neural network :math:`h_{\mathbf{\Theta}}`.
135
+ eps: (Initial) :math:`\epsilon`-value. (default: ``0.0``)
136
+ train_eps: If set to :obj:`True`, :math:`\epsilon` will be a learnable
137
+ parameter. (default: ``False``)
138
+ edge_dim: Edge feature dimensionality. (default: :obj:`None`)
139
+
140
+ Example:
141
+ ```python
142
+ import numpy as np
143
+ import keras
144
+ from k3_node.layers import GINEConv
145
+
146
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
147
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
148
+ edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
149
+
150
+ mlp = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
151
+ layer = GINEConv(mlp, edge_dim=3)
152
+ out = layer(x, edge_index, edge_attr)
153
+ print(tuple(out.shape)) # (10, 16)
154
+ ```
155
+ """
156
+
157
+ def __init__(
158
+ self,
159
+ nn: Callable,
160
+ eps: float = 0.0,
161
+ train_eps: bool = False,
162
+ edge_dim: Optional[int] = None,
163
+ **kwargs,
164
+ ):
165
+ super().__init__(aggr=kwargs.pop("aggr", kwargs.pop("aggregate", "add")), **kwargs)
166
+ self.nn = nn
167
+ self.initial_eps = eps
168
+ self.train_eps = train_eps
169
+ self.edge_dim = edge_dim
170
+
171
+ if train_eps:
172
+ self.eps = self.add_weight(
173
+ shape=(1,),
174
+ initializer=keras.initializers.Constant(eps),
175
+ name="eps",
176
+ )
177
+ else:
178
+ self.eps = ops.cast(eps, "float32")
179
+
180
+ if edge_dim is not None:
181
+ self.lin = layers.Dense(edge_dim, use_bias=True)
182
+ else:
183
+ self.lin = None
184
+
185
+ def build(self, input_shape):
186
+ if self.lin is not None:
187
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
188
+ node_dim = input_shape[0][-1]
189
+ elif isinstance(input_shape, (tuple, list)):
190
+ node_dim = input_shape[-1]
191
+ else:
192
+ node_dim = 16
193
+ self.lin.units = node_dim
194
+ self.lin.build((None, self.edge_dim))
195
+ self.built = True
196
+
197
+ def call(self, x, edge_index=None, edge_attr=None, size=None, training=None, **kwargs):
198
+ if edge_index is None and isinstance(x, (tuple, list)):
199
+ x, edge_index = x[0], x[1]
200
+
201
+ if not isinstance(x, (tuple, list)):
202
+ x_src, x_dst = x, x
203
+ else:
204
+ x_src, x_dst = x[0], x[1]
205
+
206
+ out = self.propagate(edge_index, x=(x_src, x_dst), edge_attr=edge_attr, size=size)
207
+
208
+ if x_dst is not None:
209
+ out = out + (1.0 + self.eps) * x_dst
210
+
211
+ return _apply_nn(self.nn, out, training)
212
+
213
+ def message(self, x_j, edge_attr=None):
214
+ if edge_attr is None:
215
+ return ops.relu(x_j)
216
+ if self.lin is not None:
217
+ edge_attr = self.lin(edge_attr)
218
+ return ops.relu(x_j + edge_attr)
@@ -0,0 +1,172 @@
1
+ from typing import Union, Tuple
2
+ from keras import ops
3
+ from keras.layers import Dense
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+
7
+
8
+ class GMMConv(MessagePassing):
9
+ r"""The gaussian mixture model convolutional operator from the `"Geometric
10
+ Deep Learning on Graphs and Manifolds using Mixture Model CNNs"
11
+ <https://arxiv.org/abs/1611.08402>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import GMMConv
17
+
18
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
19
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
20
+ edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
21
+
22
+ layer = GMMConv(in_channels=8, out_channels=16, dim=3, kernel_size=2)
23
+ out = layer(x, edge_index, edge_attr)
24
+ print(tuple(out.shape)) # (10, 16)
25
+ ```
26
+ """
27
+ def __init__(
28
+ self,
29
+ in_channels: Union[int, Tuple[int, int]],
30
+ out_channels: int,
31
+ dim: int,
32
+ kernel_size: int,
33
+ separate_gaussians: bool = False,
34
+ aggr: str = "mean",
35
+ root_weight: bool = True,
36
+ bias: bool = True,
37
+ **kwargs,
38
+ ):
39
+ super().__init__(aggr=aggr, **kwargs)
40
+
41
+ self.in_channels = in_channels
42
+ self.out_channels = out_channels
43
+ self.dim = dim
44
+ self.kernel_size = kernel_size
45
+ self.separate_gaussians = separate_gaussians
46
+ self.root_weight = root_weight
47
+ self.use_bias = bias
48
+
49
+ if isinstance(in_channels, int):
50
+ self.in_channels_l = in_channels
51
+ self.in_channels_r = in_channels
52
+ else:
53
+ self.in_channels_l, self.in_channels_r = in_channels
54
+
55
+ self.g = self.add_weight(
56
+ shape=(self.in_channels_l, out_channels * kernel_size),
57
+ initializer="glorot_uniform",
58
+ name="g",
59
+ )
60
+
61
+ if not separate_gaussians:
62
+ self.mu = self.add_weight(
63
+ shape=(kernel_size, dim),
64
+ initializer="glorot_uniform",
65
+ name="mu",
66
+ )
67
+ self.sigma = self.add_weight(
68
+ shape=(kernel_size, dim),
69
+ initializer="ones",
70
+ name="sigma",
71
+ )
72
+ else:
73
+ self.mu = self.add_weight(
74
+ shape=(self.in_channels_l, out_channels, kernel_size, dim),
75
+ initializer="glorot_uniform",
76
+ name="mu",
77
+ )
78
+ self.sigma = self.add_weight(
79
+ shape=(self.in_channels_l, out_channels, kernel_size, dim),
80
+ initializer="ones",
81
+ name="sigma",
82
+ )
83
+
84
+ if root_weight:
85
+ self.root = Dense(out_channels, use_bias=False)
86
+ else:
87
+ self.root = None
88
+
89
+ if bias:
90
+ self.bias = self.add_weight(
91
+ shape=(out_channels,),
92
+ initializer="zeros",
93
+ name="bias",
94
+ )
95
+ else:
96
+ self.bias = None
97
+
98
+ def build(self, input_shape=None):
99
+ if self.root is not None:
100
+ self.root.build((None, self.in_channels_r))
101
+ self.built = True
102
+
103
+ def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
104
+ if edge_index is None:
105
+ if isinstance(inputs, (list, tuple)):
106
+ if len(inputs) == 3:
107
+ x, edge_index, edge_attr = inputs
108
+ elif len(inputs) == 2:
109
+ x, edge_index = inputs
110
+ else:
111
+ raise ValueError(f"Unexpected input length {len(inputs)}")
112
+ else:
113
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
114
+ else:
115
+ x = inputs
116
+
117
+ if not self.built:
118
+ self.build()
119
+
120
+ if isinstance(x, (list, tuple)):
121
+ x_src, x_dst = x
122
+ else:
123
+ x_src = x_dst = x
124
+
125
+ num_nodes = ops.shape(x_dst)[0]
126
+
127
+ if not self.separate_gaussians:
128
+ x_l = ops.matmul(x_src, self.g)
129
+ out = self.propagate(
130
+ edge_index,
131
+ x=(x_l, x_dst),
132
+ edge_attr=edge_attr,
133
+ size=(ops.shape(x_src)[0], num_nodes),
134
+ )
135
+ else:
136
+ out = self.propagate(
137
+ edge_index,
138
+ x=(x_src, x_dst),
139
+ edge_attr=edge_attr,
140
+ size=(ops.shape(x_src)[0], num_nodes),
141
+ )
142
+
143
+ if self.root is not None:
144
+ out = out + self.root(x_dst)
145
+
146
+ if self.bias is not None:
147
+ out = out + self.bias
148
+
149
+ return out
150
+
151
+ def message(self, x_j, edge_attr):
152
+ EPS = 1e-15
153
+ M = self.out_channels
154
+ K = self.kernel_size
155
+
156
+ if not self.separate_gaussians:
157
+ # edge_attr: (E, D), mu: (K, D), sigma: (K, D)
158
+ diff = ops.expand_dims(edge_attr, 1) - ops.expand_dims(self.mu, 0)
159
+ gaussian = -0.5 * ops.square(diff) / (EPS + ops.square(ops.expand_dims(self.sigma, 0)))
160
+ gaussian = ops.exp(ops.sum(gaussian, axis=-1)) # (E, K)
161
+
162
+ x_j_reshaped = ops.reshape(x_j, (-1, K, M))
163
+ return ops.sum(x_j_reshaped * ops.expand_dims(gaussian, -1), axis=1)
164
+ else:
165
+ F = self.in_channels_l
166
+ diff = ops.expand_dims(ops.expand_dims(ops.expand_dims(edge_attr, 1), 1), 1) - ops.expand_dims(self.mu, 0)
167
+ gaussian = -0.5 * ops.square(diff) / (EPS + ops.square(ops.expand_dims(self.sigma, 0)))
168
+ gaussian = ops.exp(ops.sum(gaussian, axis=-1)) # (E, F, M, K)
169
+ g_reshaped = ops.reshape(self.g, (1, F, M, K))
170
+ gaussian = ops.sum(gaussian * g_reshaped, axis=-1) # (E, F, M)
171
+ return ops.sum(ops.expand_dims(x_j, -1) * gaussian, axis=1)
172
+