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,168 @@
1
+ import math
2
+ from typing import Optional, Union, Tuple
3
+ from keras import layers, ops
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.conv.utils import softmax
7
+ from k3_node.ops.segment import segment_sum
8
+
9
+
10
+ class TransformerConv(MessagePassing):
11
+ r"""The graph transformer operator from the `"Masked Label Prediction:
12
+ Unified Meta-Learning on Graph Neural Networks"
13
+ <https://arxiv.org/abs/2009.03509>`_ paper.
14
+
15
+ Args:
16
+ in_channels: Size of each input sample, or a tuple for bipartite graphs.
17
+ out_channels: Size of each output sample.
18
+ heads: Number of multi-head-attentions. (default: ``1``)
19
+ concat: If set to :obj:`False`, the multi-head-attentions are averaged
20
+ instead of concatenated. (default: ``True``)
21
+ beta: If set to :obj:`True`, will use a gated residual connection.
22
+ (default: ``False``)
23
+ dropout: Dropout probability of the normalized attention coefficients.
24
+ (default: ``0.0``)
25
+ edge_dim: Edge feature dimensionality (in case there are any).
26
+ (default: :obj:`None`)
27
+ bias: If set to :obj:`False`, the layer will not learn an additive bias.
28
+ (default: ``True``)
29
+ root_weight: If set to :obj:`False`, the layer will not add the
30
+ transformed root node features. (default: ``True``)
31
+
32
+ Example:
33
+ ```python
34
+ import numpy as np
35
+ from k3_node.layers import TransformerConv
36
+
37
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
38
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
39
+
40
+ layer = TransformerConv(in_channels=8, out_channels=16, heads=2)
41
+ out = layer(x, edge_index)
42
+ print(tuple(out.shape)) # (10, 32)
43
+ ```
44
+ """
45
+
46
+ def __init__(
47
+ self,
48
+ in_channels: Union[int, Tuple[int, int]],
49
+ out_channels: int,
50
+ heads: int = 1,
51
+ concat: bool = True,
52
+ beta: bool = False,
53
+ dropout: float = 0.0,
54
+ edge_dim: Optional[int] = None,
55
+ bias: bool = True,
56
+ root_weight: bool = True,
57
+ **kwargs,
58
+ ):
59
+ super().__init__(node_dim=0, aggr="add", **kwargs)
60
+ self.in_channels = in_channels
61
+ self.out_channels = out_channels
62
+ self.heads = heads
63
+ self.concat = concat
64
+ self.beta = beta and root_weight
65
+ self.root_weight = root_weight
66
+ self.dropout_rate = dropout
67
+ self.edge_dim = edge_dim
68
+ self.use_bias = bias
69
+
70
+ total_out_channels = out_channels * (heads if concat else 1)
71
+
72
+ self.lin_key = layers.Dense(heads * out_channels, use_bias=bias)
73
+ self.lin_query = layers.Dense(heads * out_channels, use_bias=bias)
74
+ self.lin_value = layers.Dense(heads * out_channels, use_bias=bias)
75
+
76
+ if edge_dim is not None:
77
+ self.lin_edge = layers.Dense(heads * out_channels, use_bias=False)
78
+ else:
79
+ self.lin_edge = None
80
+
81
+ if root_weight:
82
+ self.lin_skip = layers.Dense(total_out_channels, use_bias=bias)
83
+ if self.beta:
84
+ self.lin_beta = layers.Dense(1, use_bias=False)
85
+ else:
86
+ self.lin_beta = None
87
+ else:
88
+ self.lin_skip = None
89
+ self.lin_beta = None
90
+
91
+ self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
92
+
93
+ def build(self, input_shape):
94
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
95
+ in_channels_src = input_shape[0][-1]
96
+ in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
97
+ else:
98
+ in_channels_src = input_shape[-1]
99
+ in_channels_dst = input_shape[-1]
100
+
101
+ self.lin_key.build((None, in_channels_src))
102
+ self.lin_query.build((None, in_channels_dst))
103
+ self.lin_value.build((None, in_channels_src))
104
+
105
+ if self.lin_edge is not None:
106
+ self.lin_edge.build((None, self.edge_dim))
107
+ if self.lin_skip is not None:
108
+ self.lin_skip.build((None, in_channels_dst))
109
+ if self.lin_beta is not None:
110
+ total_out = self.out_channels * (self.heads if self.concat else 1)
111
+ self.lin_beta.build((None, 3 * total_out))
112
+
113
+ self.built = True
114
+
115
+ def call(self, x, edge_index=None, edge_attr=None, return_attention_weights=None, training=None, **kwargs):
116
+ if edge_index is None and isinstance(x, (tuple, list)):
117
+ x, edge_index = x[0], x[1]
118
+
119
+ H, C = self.heads, self.out_channels
120
+ if isinstance(x, (tuple, list)):
121
+ x_src, x_dst = x[0], x[1]
122
+ else:
123
+ x_src, x_dst = x, x
124
+
125
+ query = ops.reshape(self.lin_query(x_dst), (-1, H, C))
126
+ key = ops.reshape(self.lin_key(x_src), (-1, H, C))
127
+ value = ops.reshape(self.lin_value(x_src), (-1, H, C))
128
+
129
+ row, col = edge_index[0], edge_index[1]
130
+ row, col = ops.cast(row, "int32"), ops.cast(col, "int32")
131
+
132
+ query_i = ops.take(query, col, axis=0)
133
+ key_j = ops.take(key, row, axis=0)
134
+ value_j = ops.take(value, row, axis=0)
135
+
136
+ if self.lin_edge is not None and edge_attr is not None:
137
+ edge_attr_proj = ops.reshape(self.lin_edge(edge_attr), (-1, H, C))
138
+ key_j = key_j + edge_attr_proj
139
+ value_j = value_j + edge_attr_proj
140
+
141
+ alpha = ops.sum(query_i * key_j, axis=-1) / math.sqrt(C)
142
+ num_nodes_dst = ops.shape(x_dst)[0]
143
+ alpha = softmax(alpha, col, num_nodes=num_nodes_dst, dim=0)
144
+
145
+ if self.dropout is not None:
146
+ alpha = self.dropout(alpha, training=training)
147
+
148
+ out = ops.expand_dims(alpha, -1) * value_j
149
+ out = segment_sum(out, col, num_segments=num_nodes_dst)
150
+
151
+ if self.concat:
152
+ out = ops.reshape(out, (-1, H * C))
153
+ else:
154
+ out = ops.mean(out, axis=1)
155
+
156
+ if self.root_weight and self.lin_skip is not None:
157
+ x_r = self.lin_skip(x_dst)
158
+ if self.lin_beta is not None:
159
+ b_input = ops.concatenate([out, x_r, out - x_r], axis=-1)
160
+ b = ops.sigmoid(self.lin_beta(b_input))
161
+ out = b * x_r + (1.0 - b) * out
162
+ else:
163
+ out = out + x_r
164
+
165
+ if return_attention_weights:
166
+ return out, (edge_index, alpha)
167
+ return out
168
+
@@ -0,0 +1,403 @@
1
+ from typing import Any, Optional, Tuple, Union
2
+ from keras import ops
3
+ from k3_node.ops.segment import segment_max, segment_min, segment_sum
4
+ from k3_node.ops.creation import full
5
+
6
+
7
+ def is_tracing(x: Any) -> bool:
8
+ if x is None:
9
+ return False
10
+ try:
11
+ from keras.src.backend.common.symbolic_scope import in_symbolic_scope
12
+
13
+ if in_symbolic_scope():
14
+ return True
15
+ except Exception:
16
+ pass
17
+ name = type(x).__name__
18
+ if "Tracer" in name or "KerasTensor" in name or "SymbolicTensor" in name:
19
+ return True
20
+ if hasattr(x, "_trace"):
21
+ return True
22
+ try:
23
+ import jax
24
+ if isinstance(x, jax.core.Tracer):
25
+ return True
26
+ except Exception:
27
+ pass
28
+ try:
29
+ import tensorflow as tf
30
+ if hasattr(x, "graph") and getattr(x, "graph", None) is not None:
31
+ if not tf.executing_eagerly():
32
+ return True
33
+ except Exception:
34
+ pass
35
+ return False
36
+
37
+
38
+ def _is_compiled_trace(x) -> bool:
39
+ """True inside compiled functions, where tensor values are unknown.
40
+
41
+ On JAX, autodiff tracers created by ``jax.grad`` in eager mode (``run_eagerly=True``)
42
+ still carry concrete values, so they do not count as compiled.
43
+ """
44
+ import keras
45
+
46
+ if keras.config.backend() != "jax":
47
+ return is_tracing(x)
48
+ import jax
49
+ import jax.numpy as jnp
50
+
51
+ if not isinstance(x, jax.core.Tracer):
52
+ return False
53
+ try:
54
+ int(jnp.sum(jnp.ravel(x)[:1] * 0)) # concretizing fails only inside jit
55
+ return False
56
+ except (jax.errors.ConcretizationTypeError, jax.errors.TracerIntegerConversionError):
57
+ return True
58
+
59
+
60
+ def host_callback(fn, out_specs, *args):
61
+ """Runs the NumPy function ``fn`` on the concrete values of ``args``.
62
+
63
+ ``out_specs`` is a sequence of ``(shape, dtype)`` giving the fixed shapes of ``fn``'s outputs.
64
+ No gradient flows through the outputs. On JAX this uses ``jax.pure_callback``, which also
65
+ works under ``jax.grad``; elsewhere the arguments are converted to NumPy directly.
66
+ """
67
+ import keras
68
+ import numpy as np
69
+
70
+ def run(*values):
71
+ outs = fn(*[np.asarray(v) for v in values])
72
+ return tuple(np.asarray(o, dtype=dtype).reshape(shape) for o, (shape, dtype) in zip(outs, out_specs))
73
+
74
+ if keras.config.backend() == "jax":
75
+ import jax
76
+
77
+ specs = tuple(jax.ShapeDtypeStruct(shape, dtype) for shape, dtype in out_specs)
78
+ return jax.pure_callback(run, specs, *[jax.lax.stop_gradient(a) for a in args])
79
+ outs = run(*[ops.convert_to_numpy(ops.stop_gradient(a)) for a in args])
80
+ return tuple(ops.convert_to_tensor(o) for o in outs)
81
+
82
+
83
+ def eager_only_placeholder(layer_name: str, *tensors) -> bool:
84
+ """Guards host-side (NumPy) computations with data-dependent output sizes.
85
+
86
+ Returns ``True`` during Keras shape inference, where the caller should return a
87
+ placeholder result. Raises inside compiled functions (``tf.function`` / XLA /
88
+ ``jax.jit``), where the computation cannot run. Returns ``False`` when eager.
89
+ """
90
+ try:
91
+ from keras.src.backend.common.symbolic_scope import in_symbolic_scope
92
+
93
+ if in_symbolic_scope():
94
+ return True
95
+ except Exception:
96
+ pass
97
+ if any(_is_compiled_trace(t) for t in tensors):
98
+ raise RuntimeError(
99
+ f"{layer_name} computes a data-dependent number of clusters on the host, so it cannot run "
100
+ "inside a compiled function (tf.function, XLA or jax.jit). Compile the model with "
101
+ "`run_eagerly=True`."
102
+ )
103
+ return False
104
+
105
+
106
+ def degree(index, num_nodes: Optional[int] = None, dtype=None):
107
+ """Computes the (in/out) degree of a given index tensor.
108
+
109
+ Args:
110
+ index: 1D tensor of node indices.
111
+ num_nodes: The number of nodes.
112
+ dtype: Output data type.
113
+ """
114
+ index = ops.cast(index, "int32")
115
+ if num_nodes is None:
116
+ if is_tracing(index):
117
+ num_nodes = ops.shape(index)[0]
118
+ else:
119
+ num_nodes = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
120
+ try:
121
+ num_nodes = int(num_nodes)
122
+ except (TypeError, ValueError):
123
+ pass
124
+ ones = ops.ones((ops.shape(index)[0],), dtype=dtype or "float32")
125
+ deg = segment_sum(ones, index, num_segments=num_nodes)
126
+ if dtype is not None:
127
+ deg = ops.cast(deg, dtype)
128
+ return deg
129
+
130
+
131
+ def remove_self_loops(
132
+ edge_index,
133
+ edge_attr=None,
134
+ ) -> Tuple:
135
+ """Removes self-loops from `edge_index` and optional `edge_attr`."""
136
+ if is_tracing(edge_index):
137
+ return edge_index, edge_attr
138
+ edge_index = ops.convert_to_tensor(edge_index)
139
+ if edge_attr is not None:
140
+ edge_attr = ops.convert_to_tensor(edge_attr)
141
+ mask = edge_index[0] != edge_index[1]
142
+ where_mask = ops.where(mask)
143
+ indices = where_mask[0] if isinstance(where_mask, (list, tuple)) else where_mask
144
+ indices = ops.reshape(indices, (-1,))
145
+ edge_index = ops.take(edge_index, indices, axis=1)
146
+ if edge_attr is not None:
147
+ edge_attr = ops.take(edge_attr, indices, axis=0)
148
+ return edge_index, edge_attr
149
+
150
+
151
+ def remove_self_loops_masked(
152
+ edge_index,
153
+ edge_attr=None,
154
+ ) -> Tuple:
155
+ """Removes self-loops in a way that is also correct under static-shape tracing.
156
+
157
+ Eagerly, self-loops are dropped (as in :func:`remove_self_loops`) and the
158
+ returned mask is ``None``. Under tracing (XLA / ``jax.jit``) the edge count
159
+ must stay static, so all edges are kept and a boolean ``keep_mask`` of shape
160
+ ``[E]`` is returned that is ``False`` at the original self-loops. Callers must
161
+ exclude masked edges from aggregation.
162
+
163
+ Returns:
164
+ ``(edge_index, edge_attr, keep_mask)``
165
+ """
166
+ if not is_tracing(edge_index):
167
+ edge_index, edge_attr = remove_self_loops(edge_index, edge_attr)
168
+ return edge_index, edge_attr, None
169
+ edge_index = ops.convert_to_tensor(edge_index)
170
+ return edge_index, edge_attr, ops.not_equal(edge_index[0], edge_index[1])
171
+
172
+
173
+ def extend_mask_for_self_loops(keep_mask, num_nodes):
174
+ """Extends a ``keep_mask`` to cover the ``num_nodes`` loops appended by :func:`add_self_loops`."""
175
+ if keep_mask is None:
176
+ return None
177
+ return ops.concatenate([keep_mask, ops.ones((num_nodes,), dtype="bool")], axis=0)
178
+
179
+
180
+ def mask_edge_logits(alpha, keep_mask):
181
+ """Sets the logits of masked edges to ``-inf`` so they receive zero softmax weight."""
182
+ if keep_mask is None:
183
+ return alpha
184
+ mask = ops.reshape(keep_mask, (-1,) + (1,) * (len(alpha.shape) - 1))
185
+ return ops.where(mask, alpha, float("-inf")) # a scalar also works on torch's meta device
186
+
187
+
188
+ def add_self_loops(
189
+ edge_index,
190
+ edge_attr=None,
191
+ fill_value: Union[float, str, None] = None,
192
+ num_nodes: Optional[int] = None,
193
+ ) -> Tuple:
194
+ """Adds self-loops to `edge_index` and optional `edge_attr`."""
195
+ if is_tracing(edge_index) or is_tracing(num_nodes):
196
+ if num_nodes is None:
197
+ return edge_index, edge_attr
198
+ edge_index = ops.convert_to_tensor(edge_index)
199
+ if num_nodes is None:
200
+ num_nodes = int(ops.max(edge_index)) + 1 if ops.shape(edge_index)[1] > 0 else 0
201
+ else:
202
+ try:
203
+ num_nodes = int(num_nodes)
204
+ except (TypeError, ValueError):
205
+ pass
206
+
207
+ loop_index = ops.arange(0, num_nodes, dtype=edge_index.dtype)
208
+ loop_index = ops.stack([loop_index, loop_index], axis=0)
209
+ edge_index = ops.concatenate([edge_index, loop_index], axis=1)
210
+
211
+ if edge_attr is not None:
212
+ edge_attr = ops.convert_to_tensor(edge_attr)
213
+ attr_shape = (num_nodes,) + tuple(edge_attr.shape[1:]) if hasattr(edge_attr, "shape") else (num_nodes,)
214
+ if fill_value is None:
215
+ loop_attr = ops.zeros(attr_shape, dtype=edge_attr.dtype)
216
+ elif isinstance(fill_value, (int, float)):
217
+ loop_attr = full(attr_shape, fill_value, dtype=edge_attr.dtype)
218
+ elif fill_value == "add" or fill_value == "mean":
219
+ loop_attr = ops.zeros(attr_shape, dtype=edge_attr.dtype)
220
+ else:
221
+ loop_attr = full(attr_shape, fill_value, dtype=edge_attr.dtype)
222
+ edge_attr = ops.concatenate([edge_attr, loop_attr], axis=0)
223
+
224
+ return edge_index, edge_attr
225
+
226
+
227
+ def gcn_norm(
228
+ edge_index,
229
+ edge_weight=None,
230
+ num_nodes: Optional[int] = None,
231
+ improved: bool = False,
232
+ add_self_loops: bool = True,
233
+ flow: str = "source_to_target",
234
+ dtype=None,
235
+ ) -> Tuple:
236
+ """Computes the GCN normalization coefficients."""
237
+ fill_value = 2.0 if improved else 1.0
238
+ edge_index = ops.convert_to_tensor(edge_index)
239
+ if edge_weight is not None:
240
+ edge_weight = ops.convert_to_tensor(edge_weight)
241
+
242
+ if num_nodes is None:
243
+ if is_tracing(edge_index):
244
+ num_nodes = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
245
+ else:
246
+ num_nodes = int(ops.max(edge_index)) + 1 if ops.shape(edge_index)[1] > 0 else 0
247
+ try:
248
+ num_nodes = int(num_nodes)
249
+ except (TypeError, ValueError):
250
+ pass
251
+
252
+ if edge_weight is None:
253
+ num_edges = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
254
+ edge_weight = ops.ones((num_edges,), dtype=dtype or "float32")
255
+
256
+ if add_self_loops:
257
+ edge_index, edge_weight = globals()["add_self_loops"](
258
+ edge_index, edge_weight, fill_value=fill_value, num_nodes=num_nodes
259
+ )
260
+
261
+ row, col = edge_index[0], edge_index[1]
262
+ idx = col if flow == "source_to_target" else row
263
+ row_cast = ops.cast(row, "int32")
264
+ col_cast = ops.cast(col, "int32")
265
+ idx_cast = ops.cast(idx, "int32")
266
+
267
+ deg = segment_sum(edge_weight, idx_cast, num_segments=num_nodes)
268
+ deg_inv_sqrt = ops.power(deg, -0.5)
269
+ deg_inv_sqrt = ops.where(
270
+ ops.isinf(deg_inv_sqrt) | ops.isnan(deg_inv_sqrt), 0.0, deg_inv_sqrt
271
+ )
272
+
273
+ norm = ops.take(deg_inv_sqrt, row_cast, axis=0) * edge_weight * ops.take(deg_inv_sqrt, col_cast, axis=0)
274
+ return edge_index, norm
275
+
276
+
277
+ def get_laplacian(
278
+ edge_index,
279
+ edge_weight=None,
280
+ normalization: Optional[str] = None,
281
+ dtype=None,
282
+ num_nodes: Optional[int] = None,
283
+ ) -> Tuple:
284
+ """Computes the graph Laplacian of the given graph."""
285
+ edge_index = ops.convert_to_tensor(edge_index)
286
+ if edge_weight is not None:
287
+ edge_weight = ops.convert_to_tensor(edge_weight)
288
+
289
+ if num_nodes is None:
290
+ if is_tracing(edge_index):
291
+ num_nodes = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
292
+ else:
293
+ num_nodes = int(ops.max(edge_index)) + 1 if ops.shape(edge_index)[1] > 0 else 0
294
+ try:
295
+ num_nodes = int(num_nodes)
296
+ except (TypeError, ValueError):
297
+ pass
298
+
299
+ if edge_weight is None:
300
+ num_edges = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
301
+ edge_weight = ops.ones((num_edges,), dtype=dtype or "float32")
302
+
303
+ row, col = ops.cast(edge_index[0], "int32"), ops.cast(edge_index[1], "int32")
304
+ deg = degree(row, num_nodes=num_nodes, dtype=edge_weight.dtype)
305
+
306
+ if normalization is None:
307
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
308
+ edge_weight = ops.concatenate([-edge_weight, deg], axis=0)
309
+ elif normalization == "sym":
310
+ deg_inv_sqrt = ops.power(deg, -0.5)
311
+ deg_inv_sqrt = ops.where(
312
+ ops.isinf(deg_inv_sqrt) | ops.isnan(deg_inv_sqrt), 0.0, deg_inv_sqrt
313
+ )
314
+ edge_weight = (
315
+ ops.take(deg_inv_sqrt, row, axis=0)
316
+ * (-edge_weight)
317
+ * ops.take(deg_inv_sqrt, col, axis=0)
318
+ )
319
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
320
+ edge_weight = ops.concatenate(
321
+ [edge_weight, ops.ones((num_nodes,), dtype=edge_weight.dtype)], axis=0
322
+ )
323
+ elif normalization == "rw":
324
+ deg_inv = 1.0 / deg
325
+ deg_inv = ops.where(
326
+ ops.isinf(deg_inv) | ops.isnan(deg_inv), 0.0, deg_inv
327
+ )
328
+ edge_weight = ops.take(deg_inv, row, axis=0) * (-edge_weight)
329
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
330
+ edge_weight = ops.concatenate(
331
+ [edge_weight, ops.ones((num_nodes,), dtype=edge_weight.dtype)], axis=0
332
+ )
333
+ return edge_index, edge_weight
334
+
335
+
336
+ def _infer_dim_size(index, dim_size=None):
337
+ if dim_size is not None:
338
+ return dim_size
339
+ if hasattr(index, "is_meta") and index.is_meta:
340
+ return None
341
+ try:
342
+ if hasattr(index, "numpy") and not hasattr(index, "_has_symbolic_representation"):
343
+ # NumPy array or eager tensor with numpy()
344
+ import torch
345
+ if not isinstance(index, torch.Tensor):
346
+ return int(index.numpy().max()) + 1 if index.shape[0] > 0 else 0
347
+ if hasattr(index, "max"):
348
+ max_val = index.max()
349
+ if hasattr(max_val, "item"):
350
+ return int(max_val.item()) + 1
351
+ return int(max_val) + 1
352
+ return int(ops.convert_to_numpy(ops.max(index))) + 1 if ops.shape(index)[0] > 0 else 0
353
+ except Exception:
354
+ pass
355
+ try:
356
+ return ops.cast(ops.max(index), "int32") + 1
357
+ except Exception:
358
+ return None
359
+
360
+
361
+ def softmax(src, index, num_nodes: Optional[int] = None, dim: int = -2):
362
+ """Computes a sparsely evaluated softmax over index."""
363
+ index = ops.cast(index, "int32")
364
+ num_nodes = _infer_dim_size(index, num_nodes)
365
+
366
+ max_val = segment_max(src, index, num_segments=num_nodes)
367
+ max_val = ops.take(max_val, index, axis=dim)
368
+ exp = ops.exp(src - max_val)
369
+ sum_val = segment_sum(exp, index, num_segments=num_nodes)
370
+ sum_val = ops.take(sum_val, index, axis=dim)
371
+ return exp / (sum_val + 1e-12)
372
+
373
+
374
+ def scatter(src, index, dim=0, dim_size=None, reduce="sum"):
375
+ """Computes scatter / segment reduction."""
376
+ index = ops.cast(index, "int32")
377
+ dim_size = _infer_dim_size(index, dim_size)
378
+
379
+ if reduce in ("add", "sum"):
380
+ return segment_sum(src, index, num_segments=dim_size)
381
+ elif reduce == "mean":
382
+ sum_val = segment_sum(src, index, num_segments=dim_size)
383
+ ones = ops.ones_like(src)
384
+ count = segment_sum(ones, index, num_segments=dim_size)
385
+ count = ops.maximum(count, 1.0)
386
+ return sum_val / count
387
+ elif reduce == "max":
388
+ return segment_max(src, index, num_segments=dim_size)
389
+ val = segment_max(src, index, num_segments=dim_size)
390
+ ones = ops.ones((ops.shape(index)[0], 1), dtype=src.dtype)
391
+ count = segment_sum(ones, index, num_segments=dim_size)
392
+ return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
393
+ elif reduce == "min":
394
+ return segment_min(src, index, num_segments=dim_size)
395
+ val = segment_min(src, index, num_segments=dim_size)
396
+ ones = ops.ones((ops.shape(index)[0], 1), dtype=src.dtype)
397
+ count = segment_sum(ones, index, num_segments=dim_size)
398
+ return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
399
+ else:
400
+ raise ValueError(f"Unknown reduce operation: {reduce}")
401
+
402
+
403
+