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,228 @@
1
+ """High-level Graph / Molecular Regression Task."""
2
+
3
+ from typing import Any, Dict, List, Optional, Union
4
+ import keras
5
+ from keras import layers, ops
6
+ import numpy as np
7
+
8
+ from k3_node.tasks.base import BaseTask
9
+ from k3_node.tasks.backbone_resolver import resolve_backbone
10
+ from k3_node.layers import pool as k3_pool
11
+ from k3_node.loader import DataLoader
12
+
13
+
14
+ class GraphRegressionModel(keras.Model):
15
+ r"""Combines GNN backbone, pooling, and a continuous regression head."""
16
+
17
+ def __init__(
18
+ self,
19
+ backbone: keras.Model,
20
+ pooling: str = "add",
21
+ out_channels: int = 1,
22
+ dropout: float = 0.0,
23
+ ):
24
+ super().__init__()
25
+ self.backbone = backbone
26
+ self.pooling = pooling
27
+ self.dropout = layers.Dropout(dropout) if dropout > 0 else None
28
+ self.head = layers.Dense(out_channels)
29
+ self.num_graphs = None
30
+
31
+ def call(self, inputs, training=False):
32
+ if isinstance(inputs, (tuple, list)):
33
+ x, edge_index = inputs[0], inputs[1]
34
+ batch = inputs[2] if len(inputs) > 2 else None
35
+ size = inputs[3] if len(inputs) > 3 else self.num_graphs
36
+ else:
37
+ x = inputs
38
+ edge_index = getattr(x, "edge_index", None)
39
+ batch = getattr(x, "batch", None)
40
+ size = getattr(x, "num_graphs", self.num_graphs)
41
+ x = getattr(x, "x", x)
42
+
43
+ if batch is None:
44
+ batch = ops.zeros((ops.shape(x)[0],), dtype="int64")
45
+
46
+ h = self.backbone((x, edge_index), training=training)
47
+
48
+ if self.pooling in ("add", "sum", "global_add_pool"):
49
+ g = k3_pool.global_add_pool(h, batch, size=size)
50
+ elif self.pooling in ("mean", "global_mean_pool"):
51
+ g = k3_pool.global_mean_pool(h, batch, size=size)
52
+ elif self.pooling in ("max", "global_max_pool"):
53
+ g = k3_pool.global_max_pool(h, batch, size=size)
54
+ else:
55
+ g = k3_pool.global_add_pool(h, batch, size=size)
56
+
57
+ if self.dropout is not None:
58
+ g = self.dropout(g, training=training)
59
+ return self.head(g)
60
+
61
+
62
+ class GraphRegressor(BaseTask):
63
+ r"""High-level estimator for graph and molecular property regression tasks.
64
+
65
+ Args:
66
+ backbone: Architecture (``"schnet"``, ``"dimenet++"``, ``"attentive_fp"``,
67
+ ``"pna"``, ``"gin"``, ``"gcn"``, etc.) or custom model. (default: ``"gin"``)
68
+ in_channels (int, optional): Size of input node features.
69
+ hidden_channels (int, optional): Hidden feature dimension. (default: ``64``)
70
+ out_channels (int, optional): Number of continuous target variables. (default: ``1``)
71
+ num_layers (int, optional): Number of GNN layers. (default: ``3``)
72
+ pooling (str, optional): Readout pooling (``"add"``, ``"mean"``, ``"max"``). (default: ``"add"``)
73
+ loss (str, optional): Loss function (``"mae"`` or ``"mse"``). (default: ``"mae"``)
74
+ **backbone_kwargs: Additional arguments forwarded to the backbone constructor.
75
+ """
76
+
77
+ def __init__(
78
+ self,
79
+ backbone: Union[str, keras.Model] = "gin",
80
+ in_channels: Optional[int] = None,
81
+ hidden_channels: int = 64,
82
+ out_channels: int = 1,
83
+ num_layers: int = 3,
84
+ pooling: str = "add",
85
+ loss: str = "mae",
86
+ **backbone_kwargs,
87
+ ):
88
+ super().__init__()
89
+ self.backbone = backbone
90
+ self.in_channels = in_channels
91
+ self.hidden_channels = hidden_channels
92
+ self.out_channels = out_channels
93
+ self.num_layers = num_layers
94
+ self.pooling = pooling
95
+ self.loss_name = loss
96
+ self.backbone_kwargs = backbone_kwargs
97
+
98
+ def _init_model(self, sample_data: Any):
99
+ name = str(self.backbone).lower()
100
+ is_direct_graph_model = any(k in name for k in ("schnet", "dimenet", "attentive"))
101
+
102
+ in_c = self.in_channels or getattr(sample_data, "num_node_features", None) or getattr(sample_data, "num_features", None) or 16
103
+ self.in_channels = in_c
104
+
105
+ resolved = resolve_backbone(
106
+ self.backbone,
107
+ in_channels=in_c,
108
+ out_channels=self.out_channels,
109
+ hidden_channels=self.hidden_channels,
110
+ num_layers=self.num_layers,
111
+ **self.backbone_kwargs,
112
+ )
113
+
114
+ if is_direct_graph_model:
115
+ self.model = resolved
116
+ else:
117
+ self.model = GraphRegressionModel(
118
+ backbone=resolved,
119
+ pooling=self.pooling,
120
+ out_channels=self.out_channels,
121
+ )
122
+
123
+ def fit(
124
+ self,
125
+ dataset: Any,
126
+ epochs: int = 20,
127
+ lr: float = 0.001,
128
+ batch_size: int = 32,
129
+ shuffle: bool = True,
130
+ verbose: int = 1,
131
+ callbacks: Optional[List[Any]] = None,
132
+ ):
133
+ r"""Trains the graph regressor."""
134
+ if not isinstance(dataset, DataLoader):
135
+ loader = DataLoader(dataset, batch_size=batch_size, shuffle=shuffle)
136
+ sample = dataset[0]
137
+ else:
138
+ loader = dataset
139
+ sample = next(iter(loader))
140
+
141
+ if self.model is None:
142
+ self._init_model(sample)
143
+
144
+ if not self._is_compiled:
145
+ loss = keras.losses.MeanAbsoluteError() if self.loss_name == "mae" else keras.losses.MeanSquaredError()
146
+ self.model.compile(
147
+ optimizer=keras.optimizers.Adam(learning_rate=lr),
148
+ loss=loss,
149
+ metrics=[keras.metrics.MeanAbsoluteError(name="mae")],
150
+ )
151
+ self._is_compiled = True
152
+
153
+ history = {"loss": [], "mae": []}
154
+ for epoch in range(epochs):
155
+ batch_losses = []
156
+ batch_maes = []
157
+ for batch in loader:
158
+ x = ops.convert_to_tensor(batch.x, dtype="float32")
159
+ edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
160
+ batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
161
+ y = ops.cast(ops.convert_to_tensor(batch.y), "float32")
162
+ if len(ops.shape(y)) == 1:
163
+ y = ops.expand_dims(y, axis=-1)
164
+ num_g = int(ops.shape(y)[0])
165
+ if hasattr(self.model, "num_graphs"):
166
+ self.model.num_graphs = num_g
167
+
168
+ if not self.model.built:
169
+ y_pred = self.model((x, edge_index, batch_vec), training=False)
170
+ self.model.built = True
171
+ if hasattr(self.model, "_compile_loss") and self.model._compile_loss is not None:
172
+ self.model._compile_loss.build(y, y_pred)
173
+ if hasattr(self.model, "_compile_metrics") and self.model._compile_metrics is not None:
174
+ self.model._compile_metrics.build(y, y_pred)
175
+ if self.model.optimizer is not None and not self.model.optimizer.built:
176
+ self.model.optimizer.build(self.model.trainable_variables)
177
+
178
+ res = self.model.train_on_batch((x, edge_index, batch_vec), y)
179
+ if isinstance(res, (list, tuple)):
180
+ batch_losses.append(float(res[0]))
181
+ if len(res) > 1:
182
+ batch_maes.append(float(res[1]))
183
+ else:
184
+ batch_losses.append(float(res))
185
+
186
+ avg_loss = float(np.mean(batch_losses)) if batch_losses else 0.0
187
+ avg_mae = float(np.mean(batch_maes)) if batch_maes else 0.0
188
+ history["loss"].append(avg_loss)
189
+ history["mae"].append(avg_mae)
190
+ if verbose:
191
+ print(f"Epoch {epoch + 1}/{epochs} - loss: {avg_loss:.4f} - mae: {avg_mae:.4f}")
192
+
193
+ return history
194
+
195
+ def predict(self, dataset_or_loader: Any, batch_size: int = 32):
196
+ r"""Returns continuous predictions for graphs."""
197
+ if not isinstance(dataset_or_loader, DataLoader):
198
+ loader = DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
199
+ else:
200
+ loader = dataset_or_loader
201
+
202
+ preds = []
203
+ for batch in loader:
204
+ x = ops.convert_to_tensor(batch.x, dtype="float32")
205
+ edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
206
+ batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
207
+ num_g = int(ops.convert_to_numpy(ops.max(batch_vec))) + 1 if ops.shape(batch_vec)[0] > 0 else 1
208
+ if hasattr(self.model, "num_graphs"):
209
+ self.model.num_graphs = num_g
210
+ pred = self.model((x, edge_index, batch_vec), training=False)
211
+ preds.append(pred)
212
+ return ops.concatenate(preds, axis=0)
213
+
214
+ def evaluate(self, dataset_or_loader: Any, batch_size: int = 32) -> Dict[str, float]:
215
+ r"""Evaluates MAE on the dataset."""
216
+ preds = self.predict(dataset_or_loader, batch_size=batch_size)
217
+ ys = []
218
+ loader = dataset_or_loader if isinstance(dataset_or_loader, DataLoader) else DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
219
+ for batch in loader:
220
+ y_b = ops.cast(batch.y, "float32")
221
+ if len(ops.shape(y_b)) == 1:
222
+ y_b = ops.expand_dims(y_b, axis=-1)
223
+ ys.append(y_b)
224
+ y_all = ops.concatenate(ys, axis=0)
225
+ diff = ops.abs(preds - y_all)
226
+ mae = float(ops.convert_to_numpy(ops.mean(diff)))
227
+ mse = float(ops.convert_to_numpy(ops.mean(ops.square(diff))))
228
+ return {"mae": mae, "mse": mse, "loss": mae}
@@ -0,0 +1,306 @@
1
+ """High-level Link Prediction Task."""
2
+
3
+ from typing import Any, Dict, List, Optional, Tuple, Union
4
+ import keras
5
+ from keras import layers, ops
6
+ import numpy as np
7
+
8
+ from k3_node.tasks.base import BaseTask
9
+ from k3_node.tasks.backbone_resolver import resolve_backbone
10
+ from k3_node.models.utils import negative_sampling
11
+
12
+
13
+ class LinkPredictionModel(keras.Model):
14
+ r"""Internal neural network module bundling encoder and edge decoder."""
15
+
16
+ def __init__(
17
+ self,
18
+ encoder: keras.Model,
19
+ decoder_type: str = "inner_product",
20
+ hidden_channels: int = 64,
21
+ **kwargs,
22
+ ):
23
+ kwargs.setdefault("name", "link_prediction_model")
24
+ super().__init__(**kwargs)
25
+ self.encoder = encoder
26
+ self.decoder_type = decoder_type.lower()
27
+ if self.decoder_type == "mlp":
28
+ self.decoder_mlp = keras.Sequential([
29
+ layers.Dense(hidden_channels, activation="relu"),
30
+ layers.Dense(1),
31
+ ])
32
+ else:
33
+ self.decoder_mlp = None
34
+
35
+ def encode(self, inputs):
36
+ if hasattr(inputs, "inputs"):
37
+ inputs = inputs.inputs
38
+ return self.encoder(inputs)
39
+
40
+ def decode(self, z, edge_label_index):
41
+ edge_label_index = ops.convert_to_tensor(edge_label_index, dtype="int32")
42
+ src_idx = edge_label_index[0]
43
+ dst_idx = edge_label_index[1]
44
+ src = ops.take(z, src_idx, axis=0)
45
+ dst = ops.take(z, dst_idx, axis=0)
46
+
47
+ if self.decoder_type in ("inner_product", "dot"):
48
+ return ops.sum(src * dst, axis=-1)
49
+ elif self.decoder_type == "cosine":
50
+ src_norm = ops.sqrt(ops.maximum(ops.sum(ops.square(src), axis=-1, keepdims=True), 1e-8))
51
+ dst_norm = ops.sqrt(ops.maximum(ops.sum(ops.square(dst), axis=-1, keepdims=True), 1e-8))
52
+ return ops.sum((src / src_norm) * (dst / dst_norm), axis=-1)
53
+ elif self.decoder_type == "mlp":
54
+ feat = ops.concatenate([src, dst], axis=-1)
55
+ return ops.squeeze(self.decoder_mlp(feat), axis=-1)
56
+ else:
57
+ raise ValueError(f"Unknown decoder type '{self.decoder_type}'. Supported: 'inner_product', 'cosine', 'mlp'.")
58
+
59
+ def call(self, inputs, training=None):
60
+ r"""Executes forward pass.
61
+ inputs can be either:
62
+ - A tuple of (graph_inputs, edge_label_index)
63
+ - Just graph_inputs (in which case node embeddings z are returned)
64
+ """
65
+ if isinstance(inputs, (tuple, list)) and len(inputs) == 2 and (isinstance(inputs[1], np.ndarray) or ops.is_tensor(inputs[1])):
66
+ graph_inputs, edge_label_index = inputs
67
+ z = self.encode(graph_inputs)
68
+ return self.decode(z, edge_label_index)
69
+ else:
70
+ return self.encode(inputs)
71
+
72
+
73
+ class LinkPredictor(BaseTask):
74
+ r"""High-level estimator for link prediction tasks on graphs.
75
+
76
+ Args:
77
+ backbone: Model architecture string (``"gcn"``, ``"gat"``, ``"sage"``,
78
+ ``"gin"``, ``"pna"``, ``"mlp"``, etc.) or custom :class:`keras.Model`.
79
+ (default: ``"gcn"``)
80
+ in_channels (int, optional): Size of input node features. If not specified,
81
+ it is automatically inferred from the dataset during :meth:`fit`.
82
+ hidden_channels (int, optional): Dimensionality of hidden node features.
83
+ (default: ``64``)
84
+ out_channels (int, optional): Dimensionality of output node embeddings
85
+ used for link scoring. (default: ``64``)
86
+ num_layers (int, optional): Number of message passing layers. (default: ``2``)
87
+ decoder (str, optional): Type of edge score decoder (``"inner_product"``,
88
+ ``"cosine"``, or ``"mlp"``). (default: ``"inner_product"``)
89
+ dropout (float, optional): Dropout probability. (default: ``0.0``)
90
+ **backbone_kwargs: Additional arguments forwarded to the backbone constructor.
91
+ """
92
+
93
+ def __init__(
94
+ self,
95
+ backbone: Union[str, keras.Model] = "gcn",
96
+ in_channels: Optional[int] = None,
97
+ hidden_channels: int = 64,
98
+ out_channels: int = 64,
99
+ num_layers: int = 2,
100
+ decoder: str = "inner_product",
101
+ dropout: float = 0.0,
102
+ **backbone_kwargs,
103
+ ):
104
+ super().__init__()
105
+ self.backbone = backbone
106
+ self.in_channels = in_channels
107
+ self.hidden_channels = hidden_channels
108
+ self.out_channels = out_channels
109
+ self.num_layers = num_layers
110
+ self.decoder = decoder
111
+ self.dropout = dropout
112
+ self.backbone_kwargs = backbone_kwargs
113
+
114
+ if isinstance(backbone, keras.Model):
115
+ self.model = LinkPredictionModel(
116
+ encoder=backbone,
117
+ decoder_type=decoder,
118
+ hidden_channels=hidden_channels,
119
+ )
120
+
121
+ def _init_model(self, data: Any):
122
+ r"""Infers missing dimensions and instantiates the link prediction module."""
123
+ in_c = self.in_channels
124
+ if in_c is None:
125
+ if hasattr(data, "num_node_features") and data.num_node_features > 0:
126
+ in_c = data.num_node_features
127
+ elif hasattr(data, "num_features") and data.num_features > 0:
128
+ in_c = data.num_features
129
+ elif hasattr(data, "x") and data.x is not None:
130
+ in_c = int(ops.shape(data.x)[-1])
131
+ else:
132
+ raise ValueError("Could not automatically infer in_channels from data. Please specify in_channels.")
133
+
134
+ self.in_channels = in_c
135
+
136
+ encoder = resolve_backbone(
137
+ self.backbone,
138
+ in_channels=in_c,
139
+ out_channels=self.out_channels,
140
+ hidden_channels=self.hidden_channels,
141
+ num_layers=self.num_layers,
142
+ dropout=self.dropout,
143
+ **self.backbone_kwargs,
144
+ )
145
+
146
+ self.model = LinkPredictionModel(
147
+ encoder=encoder,
148
+ decoder_type=self.decoder,
149
+ hidden_channels=self.hidden_channels,
150
+ )
151
+
152
+ def fit(
153
+ self,
154
+ data: Any,
155
+ edge_label_index: Optional[Any] = None,
156
+ edge_label: Optional[Any] = None,
157
+ epochs: int = 20,
158
+ lr: float = 0.01,
159
+ weight_decay: float = 0.0,
160
+ neg_ratio: float = 1.0,
161
+ verbose: int = 1,
162
+ callbacks: Optional[List[Any]] = None,
163
+ ):
164
+ r"""Trains the link predictor on graph connectivity."""
165
+ if self.model is None:
166
+ self._init_model(data)
167
+
168
+ if not self._is_compiled:
169
+ opt = keras.optimizers.Adam(learning_rate=lr, weight_decay=weight_decay)
170
+ loss = keras.losses.BinaryCrossentropy(from_logits=True)
171
+ metrics = [keras.metrics.BinaryAccuracy(name="acc", threshold=0.0)]
172
+ self.model.compile(optimizer=opt, loss=loss, metrics=metrics)
173
+ self._is_compiled = True
174
+
175
+ graph_inputs = self._extract_inputs(data)
176
+
177
+ # Check for pre-split edge labels on data or passed explicitly
178
+ has_labels = edge_label_index is not None and edge_label is not None
179
+ if not has_labels:
180
+ if hasattr(data, "train_edge_label_index") and hasattr(data, "train_edge_label"):
181
+ edge_label_index = data.train_edge_label_index
182
+ edge_label = data.train_edge_label
183
+ has_labels = True
184
+ elif hasattr(data, "edge_label_index") and hasattr(data, "edge_label"):
185
+ edge_label_index = data.edge_label_index
186
+ edge_label = data.edge_label
187
+ has_labels = True
188
+
189
+ num_nodes = None
190
+ if hasattr(data, "num_nodes") and data.num_nodes is not None:
191
+ num_nodes = data.num_nodes
192
+ elif hasattr(data, "x") and data.x is not None:
193
+ num_nodes = int(ops.shape(data.x)[0])
194
+
195
+ if not has_labels:
196
+ pos_edge_index = data.edge_index
197
+ pos_np = ops.convert_to_numpy(pos_edge_index).astype(np.int32)
198
+ num_pos = pos_np.shape[1]
199
+ num_neg = int(num_pos * neg_ratio)
200
+
201
+ history = {"loss": [], "acc": []}
202
+ for epoch in range(epochs):
203
+ if has_labels:
204
+ total_edges = ops.convert_to_tensor(edge_label_index, dtype="int32")
205
+ labels = ops.cast(ops.convert_to_tensor(edge_label), "float32")
206
+ else:
207
+ neg_np = negative_sampling(pos_np, num_nodes=num_nodes, num_neg_samples=num_neg)
208
+ total_edges = ops.convert_to_tensor(np.concatenate([pos_np, neg_np], axis=1), dtype="int32")
209
+ labels = ops.convert_to_tensor(
210
+ np.concatenate([np.ones(num_pos, dtype=np.float32), np.zeros(num_neg, dtype=np.float32)]),
211
+ dtype="float32",
212
+ )
213
+
214
+ if not self.model.built:
215
+ y_pred = self.model((graph_inputs, total_edges), training=False)
216
+ self.model.built = True
217
+ if hasattr(self.model, "_compile_loss") and self.model._compile_loss is not None:
218
+ self.model._compile_loss.build(labels, y_pred)
219
+ if hasattr(self.model, "_compile_metrics") and self.model._compile_metrics is not None:
220
+ self.model._compile_metrics.build(labels, y_pred)
221
+ if self.model.optimizer is not None and not self.model.optimizer.built:
222
+ self.model.optimizer.build(self.model.trainable_variables)
223
+
224
+ res = self.model.train_on_batch((graph_inputs, total_edges), labels)
225
+ if isinstance(res, (list, tuple)):
226
+ l, a = float(res[0]), float(res[1]) if len(res) > 1 else 0.0
227
+ else:
228
+ l, a = float(res), 0.0
229
+ history["loss"].append(l)
230
+ history["acc"].append(a)
231
+ if verbose:
232
+ print(f"Epoch {epoch + 1}/{epochs} - loss: {l:.4f} - acc: {a:.4f}")
233
+
234
+ return history
235
+
236
+ def encode(self, data: Any):
237
+ r"""Computes latent node representations for the input graph."""
238
+ if self.model is None:
239
+ raise RuntimeError("Model is not initialized. Call fit() or construct with a model first.")
240
+ graph_inputs = self._extract_inputs(data)
241
+ return self.model.encode(graph_inputs)
242
+
243
+ def predict_proba(self, data: Any, edge_label_index: Optional[Any] = None):
244
+ r"""Predicts link existence probabilities for edge pairs."""
245
+ if self.model is None:
246
+ raise RuntimeError("Model is not initialized. Fit or load a model first.")
247
+
248
+ if edge_label_index is None:
249
+ if hasattr(data, "test_edge_label_index"):
250
+ edge_label_index = data.test_edge_label_index
251
+ elif hasattr(data, "edge_label_index"):
252
+ edge_label_index = data.edge_label_index
253
+ elif hasattr(data, "edge_index"):
254
+ edge_label_index = data.edge_index
255
+ else:
256
+ raise ValueError("No edge_label_index provided and none found on data object.")
257
+
258
+ z = self.encode(data)
259
+ logits = self.model.decode(z, edge_label_index)
260
+ return ops.sigmoid(logits)
261
+
262
+ def predict(self, data: Any, edge_label_index: Optional[Any] = None, threshold: float = 0.5):
263
+ r"""Predicts binary link presence (0 or 1) for edge pairs."""
264
+ probs = self.predict_proba(data, edge_label_index=edge_label_index)
265
+ return ops.cast(probs >= threshold, "int64")
266
+
267
+ def evaluate(
268
+ self,
269
+ data: Any,
270
+ edge_label_index: Optional[Any] = None,
271
+ edge_label: Optional[Any] = None,
272
+ ) -> Dict[str, float]:
273
+ r"""Evaluates link prediction performance (AUC, AP, Accuracy)."""
274
+ if edge_label_index is None and edge_label is None:
275
+ if hasattr(data, "test_edge_label_index") and hasattr(data, "test_edge_label"):
276
+ edge_label_index = data.test_edge_label_index
277
+ edge_label = data.test_edge_label
278
+ elif hasattr(data, "edge_label_index") and hasattr(data, "edge_label"):
279
+ edge_label_index = data.edge_label_index
280
+ edge_label = data.edge_label
281
+ else:
282
+ # Sample negative edges against edge_index
283
+ pos_edges = ops.convert_to_numpy(data.edge_index).astype(np.int32)
284
+ num_nodes = data.num_nodes if hasattr(data, "num_nodes") else int(ops.shape(data.x)[0])
285
+ neg_edges = negative_sampling(pos_edges, num_nodes=num_nodes, num_neg_samples=pos_edges.shape[1])
286
+ edge_label_index = np.concatenate([pos_edges, neg_edges], axis=1)
287
+ edge_label = np.concatenate([np.ones(pos_edges.shape[1]), np.zeros(neg_edges.shape[1])])
288
+
289
+ probs = self.predict_proba(data, edge_label_index=edge_label_index)
290
+ probs_np = ops.convert_to_numpy(probs).flatten()
291
+ y_np = ops.convert_to_numpy(edge_label).flatten()
292
+
293
+ metrics = {}
294
+ # Accuracy
295
+ preds_bin = (probs_np >= 0.5).astype(np.float32)
296
+ metrics["accuracy"] = float(np.mean(preds_bin == y_np))
297
+
298
+ # AUC and AP via scikit-learn when available
299
+ try:
300
+ from sklearn.metrics import roc_auc_score, average_precision_score
301
+ metrics["auc"] = float(roc_auc_score(y_np, probs_np))
302
+ metrics["ap"] = float(average_precision_score(y_np, probs_np))
303
+ except ImportError:
304
+ pass
305
+
306
+ return metrics