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,244 @@
1
+ from typing import Optional, Union, Tuple
2
+ from keras import layers, 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
+ from k3_node.ops.segment import segment_sum
13
+
14
+
15
+ class GATConv(MessagePassing):
16
+ r"""The graph attentional operator from the `"Graph Attention Networks"
17
+ <https://arxiv.org/abs/1710.10903>`_ paper.
18
+
19
+ Args:
20
+ in_channels: Size of each input sample, or a tuple for bipartite graphs.
21
+ out_channels: Size of each output sample.
22
+ heads: Number of multi-head-attentions. (default: ``1``)
23
+ concat: If set to :obj:`False`, the multi-head-attentions are averaged
24
+ instead of concatenated. (default: ``True``)
25
+ negative_slope: LeakyReLU angle of the negative slope. (default: ``0.2``)
26
+ dropout: Dropout probability of the normalized attention coefficients.
27
+ (default: ``0.0``)
28
+ add_self_loops: If set to :obj:`False`, will not add self-loops to
29
+ the input graph. (default: ``True``)
30
+ edge_dim: Edge feature dimensionality (in case there are any).
31
+ (default: :obj:`None`)
32
+ fill_value: The way to generate edge features of self-loops
33
+ (default: ``"mean"``)
34
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
35
+ (default: ``True``)
36
+ share_weights: If set to :obj:`True`, the same matrix will be applied
37
+ to the source and target node features. (default: ``False``)
38
+ residual: If set to :obj:`True`, will compute residual connections.
39
+ (default: ``False``)
40
+
41
+ Example:
42
+ ```python
43
+ import numpy as np
44
+ from k3_node.layers import GATConv
45
+
46
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
47
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
48
+
49
+ layer = GATConv(in_channels=8, out_channels=16, heads=2)
50
+ out = layer(x, edge_index)
51
+ print(tuple(out.shape)) # (10, 32)
52
+ ```
53
+ """
54
+
55
+ def __init__(
56
+ self,
57
+ in_channels: Union[int, Tuple[int, int]],
58
+ out_channels: int,
59
+ heads: int = 1,
60
+ concat: bool = True,
61
+ negative_slope: float = 0.2,
62
+ dropout: float = 0.0,
63
+ add_self_loops: bool = True,
64
+ edge_dim: Optional[int] = None,
65
+ fill_value: Union[float, str] = "mean",
66
+ bias: bool = True,
67
+ share_weights: bool = False,
68
+ residual: bool = False,
69
+ **kwargs,
70
+ ):
71
+ super().__init__(node_dim=0, **kwargs)
72
+ self.in_channels = in_channels
73
+ self.out_channels = out_channels
74
+ self.heads = heads
75
+ self.concat = concat
76
+ self.negative_slope = negative_slope
77
+ self.dropout_rate = dropout
78
+ self.add_self_loops = add_self_loops
79
+ self.edge_dim = edge_dim
80
+ self.fill_value = fill_value
81
+ self.use_bias = bias
82
+ self.share_weights = share_weights
83
+ self.residual = residual
84
+
85
+ total_out_channels = out_channels * (heads if concat else 1)
86
+
87
+ if isinstance(in_channels, int):
88
+ self.lin = layers.Dense(heads * out_channels, use_bias=False)
89
+ self.lin_src = self.lin
90
+ self.lin_dst = self.lin
91
+ else:
92
+ self.lin = None
93
+ self.lin_src = layers.Dense(heads * out_channels, use_bias=False)
94
+ if share_weights:
95
+ self.lin_dst = self.lin_src
96
+ else:
97
+ self.lin_dst = layers.Dense(heads * out_channels, use_bias=False)
98
+
99
+ if edge_dim is not None:
100
+ self.lin_edge = layers.Dense(heads * out_channels, use_bias=False)
101
+ else:
102
+ self.lin_edge = None
103
+
104
+ if residual:
105
+ self.res = layers.Dense(total_out_channels, use_bias=False)
106
+ else:
107
+ self.res = None
108
+
109
+ self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
110
+
111
+ def build(self, input_shape):
112
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
113
+ in_channels_src = input_shape[0][-1]
114
+ in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
115
+ else:
116
+ in_channels_src = input_shape[-1]
117
+ in_channels_dst = input_shape[-1]
118
+
119
+ self.lin_src.build((None, in_channels_src))
120
+ if self.lin_dst is not self.lin_src:
121
+ self.lin_dst.build((None, in_channels_dst))
122
+ if self.lin_edge is not None:
123
+ self.lin_edge.build((None, self.edge_dim))
124
+ if self.res is not None:
125
+ self.res.build((None, in_channels_dst))
126
+
127
+ self.att_src = self.add_weight(
128
+ shape=(1, self.heads, self.out_channels),
129
+ initializer="glorot_uniform",
130
+ name="att_src",
131
+ )
132
+ self.att_dst = self.add_weight(
133
+ shape=(1, self.heads, self.out_channels),
134
+ initializer="glorot_uniform",
135
+ name="att_dst",
136
+ )
137
+ if self.edge_dim is not None:
138
+ self.att_edge = self.add_weight(
139
+ shape=(1, self.heads, self.out_channels),
140
+ initializer="glorot_uniform",
141
+ name="att_edge",
142
+ )
143
+ else:
144
+ self.att_edge = None
145
+
146
+ total_out_channels = self.out_channels * (self.heads if self.concat else 1)
147
+ if self.use_bias:
148
+ self.bias = self.add_weight(
149
+ shape=(total_out_channels,),
150
+ initializer="zeros",
151
+ name="bias",
152
+ )
153
+ else:
154
+ self.bias = None
155
+ self.built = True
156
+
157
+ def call(self, x, edge_index=None, edge_attr=None, size=None, return_attention_weights=None, training=None, **kwargs):
158
+ if edge_index is None and isinstance(x, (tuple, list)):
159
+ x, edge_index = x[0], x[1]
160
+
161
+ H, C = self.heads, self.out_channels
162
+ if isinstance(x, (tuple, list)):
163
+ x_src, x_dst = x[0], x[1]
164
+ else:
165
+ x_src, x_dst = x, x
166
+
167
+ x_src_proj = ops.reshape(self.lin_src(x_src), (-1, H, C))
168
+ x_dst_proj = ops.reshape(self.lin_dst(x_dst), (-1, H, C)) if x_dst is not None else None
169
+
170
+ alpha_src = ops.sum(x_src_proj * self.att_src, axis=-1)
171
+ alpha_dst = ops.sum(x_dst_proj * self.att_dst, axis=-1) if x_dst_proj is not None else None
172
+
173
+ if self.add_self_loops:
174
+ if not isinstance(x, (tuple, list)):
175
+ num_nodes = ops.shape(x)[0]
176
+ else:
177
+ num_nodes = ops.shape(x_src)[0]
178
+ if x_dst is not None:
179
+ num_nodes = ops.minimum(num_nodes, ops.shape(x_dst)[0])
180
+ edge_index, edge_attr, keep_mask = remove_self_loops_masked(edge_index, edge_attr)
181
+ edge_index, edge_attr = add_self_loops(
182
+ edge_index, edge_attr, fill_value=self.fill_value, num_nodes=num_nodes
183
+ )
184
+ keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
185
+ else:
186
+ keep_mask = None
187
+
188
+ row, col = edge_index[0], edge_index[1]
189
+ row, col = ops.cast(row, "int32"), ops.cast(col, "int32")
190
+
191
+ alpha_j = ops.take(alpha_src, row, axis=0)
192
+ alpha_i = ops.take(alpha_dst, col, axis=0) if alpha_dst is not None else 0.0
193
+ alpha = alpha_j + alpha_i
194
+
195
+ if edge_attr is not None and self.lin_edge is not None and self.att_edge is not None:
196
+ edge_attr_proj = ops.reshape(self.lin_edge(edge_attr), (-1, H, C))
197
+ alpha = alpha + ops.sum(edge_attr_proj * self.att_edge, axis=-1)
198
+
199
+ alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
200
+ alpha = mask_edge_logits(alpha, keep_mask)
201
+ num_nodes_dst = ops.shape(x_dst)[0] if x_dst is not None else ops.shape(x_src)[0]
202
+ alpha = softmax(alpha, col, num_nodes=num_nodes_dst, dim=0)
203
+
204
+ if self.dropout is not None:
205
+ alpha = self.dropout(alpha, training=training)
206
+
207
+ # Message & aggregate
208
+ x_src_j = ops.take(x_src_proj, row, axis=0)
209
+ out = ops.expand_dims(alpha, -1) * x_src_j
210
+ out = segment_sum(out, col, num_segments=num_nodes_dst)
211
+
212
+ if self.concat:
213
+ out = ops.reshape(out, (-1, H * C))
214
+ else:
215
+ out = ops.mean(out, axis=1)
216
+
217
+ if self.res is not None and x_dst is not None:
218
+ out = out + self.res(x_dst)
219
+
220
+ if self.bias is not None:
221
+ out = out + self.bias
222
+
223
+ if return_attention_weights:
224
+ return out, (edge_index, alpha)
225
+ return out
226
+
227
+
228
+ class FusedGATConv(GATConv):
229
+ r"""The fused graph attentional operator.
230
+
231
+ Example:
232
+ ```python
233
+ import numpy as np
234
+ from k3_node.layers import FusedGATConv
235
+
236
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
237
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
238
+
239
+ layer = FusedGATConv(in_channels=8, out_channels=16)
240
+ out = layer(x, edge_index)
241
+ print(tuple(out.shape)) # (10, 16)
242
+ ```
243
+ """
244
+ pass
@@ -0,0 +1,136 @@
1
+ # ported from spektral
2
+
3
+ from keras import ops
4
+ from keras.layers import GRUCell
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+
8
+
9
+ class GatedGraphConv(MessagePassing):
10
+ """
11
+ `k3_node.layers.GatedGraphConv`
12
+
13
+ Implementation of Gated Graph Convolution (GGC) layer
14
+
15
+ Args:
16
+ channels: The number of output channels.
17
+ n_layers: The number of GGC layers to stack.
18
+ activation: Activation function to use.
19
+ use_bias: Whether to add a bias to the linear transformation.
20
+ kernel_initializer: Initializer for the `kernel` weights matrix.
21
+ bias_initializer: Initializer for the bias vector.
22
+ kernel_regularizer: Regularizer for the `kernel` weights matrix.
23
+ bias_regularizer: Regularizer for the bias vector.
24
+ activity_regularizer: Regularizer for the output.
25
+ kernel_constraint: Constraint for the `kernel` weights matrix.
26
+ bias_constraint: Constraint for the bias vector.
27
+ **kwargs: Additional arguments to pass to the `MessagePassing` superclass.
28
+
29
+ Example:
30
+ ```python
31
+ import numpy as np
32
+ from k3_node.layers import GatedGraphConv
33
+
34
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
35
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
36
+
37
+ layer = GatedGraphConv(out_channels=16, num_layers=2)
38
+ out = layer(x, edge_index)
39
+ print(tuple(out.shape)) # (10, 16)
40
+ ```
41
+ """
42
+ def __init__(
43
+ self,
44
+ channels=None,
45
+ n_layers=None,
46
+ out_channels=None,
47
+ num_layers=None,
48
+ activation=None,
49
+ use_bias=True,
50
+ kernel_initializer="glorot_uniform",
51
+ bias_initializer="zeros",
52
+ kernel_regularizer=None,
53
+ bias_regularizer=None,
54
+ activity_regularizer=None,
55
+ kernel_constraint=None,
56
+ bias_constraint=None,
57
+ **kwargs,
58
+ ):
59
+ channels = out_channels if out_channels is not None else channels
60
+ n_layers = num_layers if num_layers is not None else n_layers
61
+ super().__init__(
62
+ activation=activation,
63
+ use_bias=use_bias,
64
+ kernel_initializer=kernel_initializer,
65
+ bias_initializer=bias_initializer,
66
+ kernel_regularizer=kernel_regularizer,
67
+ bias_regularizer=bias_regularizer,
68
+ activity_regularizer=activity_regularizer,
69
+ kernel_constraint=kernel_constraint,
70
+ bias_constraint=bias_constraint,
71
+ **kwargs,
72
+ )
73
+ self.channels = channels
74
+ self.out_channels = channels
75
+ self.n_layers = n_layers
76
+ self.num_layers = n_layers
77
+
78
+ def build(self, input_shape=None):
79
+ self.kernel = self.add_weight(
80
+ name="kernel",
81
+ shape=(self.n_layers, self.channels, self.channels),
82
+ initializer=self.kernel_initializer,
83
+ regularizer=self.kernel_regularizer,
84
+ constraint=self.kernel_constraint,
85
+ )
86
+ self.rnn = GRUCell(
87
+ self.channels,
88
+ kernel_initializer=self.kernel_initializer,
89
+ bias_initializer=self.bias_initializer,
90
+ kernel_regularizer=self.kernel_regularizer,
91
+ bias_regularizer=self.bias_regularizer,
92
+ activity_regularizer=self.activity_regularizer,
93
+ kernel_constraint=self.kernel_constraint,
94
+ bias_constraint=self.bias_constraint,
95
+ use_bias=self.use_bias,
96
+ dtype=self.dtype,
97
+ )
98
+ self.rnn.build((self.channels,))
99
+ super().build(input_shape)
100
+ self.built = True
101
+
102
+ def call(self, x, edge_index=None, edge_weight=None, **kwargs):
103
+ is_legacy = False
104
+ if edge_index is None and isinstance(x, (tuple, list)):
105
+ x, a, _ = self.get_inputs(x)
106
+ edge_index = a
107
+ is_legacy = True
108
+
109
+ F = ops.shape(x)[-1]
110
+ if F < self.channels:
111
+ to_pad = self.channels - F
112
+ ndims = len(ops.shape(x)) - 1
113
+ output = ops.pad(x, [[0, 0]] * ndims + [[0, to_pad]])
114
+ elif F > self.channels:
115
+ output = x[..., :self.channels]
116
+ else:
117
+ output = x
118
+
119
+ for i in range(self.n_layers):
120
+ m = ops.matmul(output, self.kernel[i])
121
+ if is_legacy:
122
+ m = self.propagate(m, edge_index)
123
+ else:
124
+ m = self.propagate(edge_index, x=m, edge_weight=edge_weight)
125
+ output = self.rnn(m, [output])[0]
126
+
127
+ if hasattr(self, "activation") and self.activation is not None and callable(self.activation):
128
+ output = self.activation(output)
129
+ return output
130
+
131
+ @property
132
+ def config(self):
133
+ return {
134
+ "channels": self.channels,
135
+ "n_layers": self.n_layers,
136
+ }
@@ -0,0 +1,205 @@
1
+ from typing import Optional, Union, Tuple
2
+ from keras import layers, 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
+ from k3_node.ops.segment import segment_sum
13
+
14
+
15
+ class GATv2Conv(MessagePassing):
16
+ r"""The GATv2 operator from the `"How Attentive are Graph Attention Networks?"
17
+ <https://arxiv.org/abs/2105.14491>`_ paper, which fixes the static
18
+ attention problem of standard :class:`~k3_node.layers.conv.GATConv`.
19
+
20
+ Args:
21
+ in_channels: Size of each input sample, or a tuple for bipartite graphs.
22
+ out_channels: Size of each output sample.
23
+ heads: Number of multi-head-attentions. (default: ``1``)
24
+ concat: If set to :obj:`False`, the multi-head-attentions are averaged
25
+ instead of concatenated. (default: ``True``)
26
+ negative_slope: LeakyReLU angle of the negative slope. (default: ``0.2``)
27
+ dropout: Dropout probability of the normalized attention coefficients.
28
+ (default: ``0.0``)
29
+ add_self_loops: If set to :obj:`False`, will not add self-loops to
30
+ the input graph. (default: ``True``)
31
+ edge_dim: Edge feature dimensionality (in case there are any).
32
+ (default: :obj:`None`)
33
+ fill_value: The way to generate edge features of self-loops
34
+ (default: ``"mean"``)
35
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
36
+ (default: ``True``)
37
+ share_weights: If set to :obj:`True`, the same matrix will be applied
38
+ to the source and target node features. (default: ``False``)
39
+ residual: If set to :obj:`True`, will compute residual connections.
40
+ (default: ``False``)
41
+
42
+ Example:
43
+ ```python
44
+ import numpy as np
45
+ from k3_node.layers import GATv2Conv
46
+
47
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
48
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
49
+
50
+ layer = GATv2Conv(in_channels=8, out_channels=16, heads=2)
51
+ out = layer(x, edge_index)
52
+ print(tuple(out.shape)) # (10, 32)
53
+ ```
54
+ """
55
+
56
+ def __init__(
57
+ self,
58
+ in_channels: Union[int, Tuple[int, int]],
59
+ out_channels: int,
60
+ heads: int = 1,
61
+ concat: bool = True,
62
+ negative_slope: float = 0.2,
63
+ dropout: float = 0.0,
64
+ add_self_loops: bool = True,
65
+ edge_dim: Optional[int] = None,
66
+ fill_value: Union[float, str] = "mean",
67
+ bias: bool = True,
68
+ share_weights: bool = False,
69
+ residual: bool = False,
70
+ **kwargs,
71
+ ):
72
+ super().__init__(node_dim=0, **kwargs)
73
+ self.in_channels = in_channels
74
+ self.out_channels = out_channels
75
+ self.heads = heads
76
+ self.concat = concat
77
+ self.negative_slope = negative_slope
78
+ self.dropout_rate = dropout
79
+ self.add_self_loops = add_self_loops
80
+ self.edge_dim = edge_dim
81
+ self.fill_value = fill_value
82
+ self.use_bias = bias
83
+ self.share_weights = share_weights
84
+ self.residual = residual
85
+
86
+ total_out_channels = out_channels * (heads if concat else 1)
87
+
88
+ self.lin_l = layers.Dense(heads * out_channels, use_bias=bias)
89
+ if share_weights:
90
+ self.lin_r = self.lin_l
91
+ else:
92
+ self.lin_r = layers.Dense(heads * out_channels, use_bias=bias)
93
+
94
+ if edge_dim is not None:
95
+ self.lin_edge = layers.Dense(heads * out_channels, use_bias=False)
96
+ else:
97
+ self.lin_edge = None
98
+
99
+ if residual:
100
+ self.res = layers.Dense(total_out_channels, use_bias=False)
101
+ else:
102
+ self.res = None
103
+
104
+ self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
105
+
106
+ def build(self, input_shape):
107
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
108
+ in_channels_src = input_shape[0][-1]
109
+ in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
110
+ else:
111
+ in_channels_src = input_shape[-1]
112
+ in_channels_dst = input_shape[-1]
113
+
114
+ self.lin_l.build((None, in_channels_src))
115
+ if self.lin_r is not self.lin_l:
116
+ self.lin_r.build((None, in_channels_dst))
117
+ if self.lin_edge is not None:
118
+ self.lin_edge.build((None, self.edge_dim))
119
+ if self.res is not None:
120
+ self.res.build((None, in_channels_dst))
121
+
122
+ self.att = self.add_weight(
123
+ shape=(1, self.heads, self.out_channels),
124
+ initializer="glorot_uniform",
125
+ name="att",
126
+ )
127
+
128
+ total_out_channels = self.out_channels * (self.heads if self.concat else 1)
129
+ if self.use_bias:
130
+ self.bias = self.add_weight(
131
+ shape=(total_out_channels,),
132
+ initializer="zeros",
133
+ name="bias",
134
+ )
135
+ else:
136
+ self.bias = None
137
+ self.built = True
138
+
139
+ def call(self, x, edge_index=None, edge_attr=None, size=None, return_attention_weights=None, training=None, **kwargs):
140
+ if edge_index is None and isinstance(x, (tuple, list)):
141
+ x, edge_index = x[0], x[1]
142
+
143
+ H, C = self.heads, self.out_channels
144
+ if isinstance(x, (tuple, list)):
145
+ x_src, x_dst = x[0], x[1]
146
+ else:
147
+ x_src, x_dst = x, x
148
+
149
+ x_l = ops.reshape(self.lin_l(x_src), (-1, H, C))
150
+ x_r = ops.reshape(self.lin_r(x_dst), (-1, H, C)) if x_dst is not None else x_l
151
+
152
+ if self.add_self_loops:
153
+ if not isinstance(x, (tuple, list)):
154
+ num_nodes = ops.shape(x)[0]
155
+ else:
156
+ num_nodes = ops.shape(x_l)[0]
157
+ if x_r is not None:
158
+ num_nodes = ops.minimum(num_nodes, ops.shape(x_r)[0])
159
+ edge_index, edge_attr, keep_mask = remove_self_loops_masked(edge_index, edge_attr)
160
+ edge_index, edge_attr = add_self_loops(
161
+ edge_index, edge_attr, fill_value=self.fill_value, num_nodes=num_nodes
162
+ )
163
+ keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
164
+ else:
165
+ keep_mask = None
166
+
167
+ row, col = edge_index[0], edge_index[1]
168
+ row, col = ops.cast(row, "int32"), ops.cast(col, "int32")
169
+
170
+ x_l_j = ops.take(x_l, row, axis=0)
171
+ x_r_i = ops.take(x_r, col, axis=0)
172
+ alpha = x_l_j + x_r_i
173
+
174
+ if edge_attr is not None and self.lin_edge is not None:
175
+ edge_attr_proj = ops.reshape(self.lin_edge(edge_attr), (-1, H, C))
176
+ alpha = alpha + edge_attr_proj
177
+
178
+ alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
179
+ alpha = ops.sum(alpha * self.att, axis=-1)
180
+ alpha = mask_edge_logits(alpha, keep_mask)
181
+
182
+ num_nodes_dst = ops.shape(x_r)[0]
183
+ alpha = softmax(alpha, col, num_nodes=num_nodes_dst, dim=0)
184
+
185
+ if self.dropout is not None:
186
+ alpha = self.dropout(alpha, training=training)
187
+
188
+ out = ops.expand_dims(alpha, -1) * x_l_j
189
+ out = segment_sum(out, col, num_segments=num_nodes_dst)
190
+
191
+ if self.concat:
192
+ out = ops.reshape(out, (-1, H * C))
193
+ else:
194
+ out = ops.mean(out, axis=1)
195
+
196
+ if self.res is not None and x_dst is not None:
197
+ out = out + self.res(x_dst)
198
+
199
+ if self.bias is not None:
200
+ out = out + self.bias
201
+
202
+ if return_attention_weights:
203
+ return out, (edge_index, alpha)
204
+ return out
205
+