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,194 @@
1
+ """High-level Node Classification Task."""
2
+
3
+ from typing import Any, Dict, List, Optional, Union
4
+ import keras
5
+ from keras import ops
6
+
7
+ from k3_node.tasks.base import BaseTask
8
+ from k3_node.tasks.backbone_resolver import resolve_backbone
9
+
10
+
11
+ class NodeClassifier(BaseTask):
12
+ r"""High-level estimator for node classification tasks.
13
+
14
+ Args:
15
+ backbone: Model architecture string (``"gcn"``, ``"gat"``, ``"sage"``,
16
+ ``"gin"``, ``"pna"``, ``"mlp"``, etc.) or a custom :class:`keras.Model`.
17
+ (default: ``"gcn"``)
18
+ in_channels (int, optional): Size of input node features. If not specified,
19
+ it is automatically inferred from the dataset during :meth:`fit`.
20
+ hidden_channels (int, optional): Dimensionality of hidden node features.
21
+ (default: ``64``)
22
+ out_channels (int, optional): Number of target classes. If not specified,
23
+ it is automatically inferred from the dataset during :meth:`fit`.
24
+ num_layers (int, optional): Number of message passing layers. (default: ``2``)
25
+ dropout (float, optional): Dropout probability. (default: ``0.5``)
26
+ multi_label (bool, optional): If :obj:`True`, treats the problem as multi-label
27
+ binary classification using binary crossentropy. (default: ``False``)
28
+ **backbone_kwargs: Additional arguments forwarded to the backbone constructor.
29
+ """
30
+
31
+ def __init__(
32
+ self,
33
+ backbone: Union[str, keras.Model] = "gcn",
34
+ in_channels: Optional[int] = None,
35
+ hidden_channels: int = 64,
36
+ out_channels: Optional[int] = None,
37
+ num_classes: Optional[int] = None,
38
+ num_layers: int = 2,
39
+ dropout: float = 0.5,
40
+ multi_label: bool = False,
41
+ **backbone_kwargs,
42
+ ):
43
+ super().__init__()
44
+ self.backbone = backbone
45
+ self.in_channels = in_channels
46
+ self.hidden_channels = hidden_channels
47
+ self.out_channels = out_channels if out_channels is not None else num_classes
48
+ self.num_layers = num_layers
49
+ self.dropout = dropout
50
+ self.multi_label = multi_label
51
+ self.backbone_kwargs = backbone_kwargs
52
+
53
+ if isinstance(backbone, keras.Model):
54
+ self.model = backbone
55
+
56
+ def _init_model(self, data: Any):
57
+ r"""Infers missing dimensions and instantiates the backbone model."""
58
+ in_c = self.in_channels
59
+ if in_c is None:
60
+ if hasattr(data, "num_node_features") and data.num_node_features > 0:
61
+ in_c = data.num_node_features
62
+ elif hasattr(data, "num_features") and data.num_features > 0:
63
+ in_c = data.num_features
64
+ elif hasattr(data, "x") and data.x is not None:
65
+ in_c = int(ops.shape(data.x)[-1])
66
+ else:
67
+ raise ValueError("Could not automatically infer in_channels from data. Please specify in_channels.")
68
+
69
+ out_c = self.out_channels
70
+ if out_c is None:
71
+ if hasattr(data, "num_classes") and data.num_classes is not None:
72
+ out_c = data.num_classes
73
+ elif hasattr(data, "y") and data.y is not None:
74
+ y = data.y
75
+ if self.multi_label:
76
+ out_c = int(ops.shape(y)[-1])
77
+ else:
78
+ out_c = int(ops.convert_to_numpy(ops.max(y))) + 1
79
+ else:
80
+ raise ValueError("Could not automatically infer out_channels from data. Please specify out_channels.")
81
+
82
+ if not self.multi_label:
83
+ out_c = max(int(out_c), 2)
84
+
85
+ self.in_channels = in_c
86
+ self.out_channels = out_c
87
+
88
+ self.model = resolve_backbone(
89
+ self.backbone,
90
+ in_channels=in_c,
91
+ out_channels=out_c,
92
+ hidden_channels=self.hidden_channels,
93
+ num_layers=self.num_layers,
94
+ dropout=self.dropout,
95
+ **self.backbone_kwargs,
96
+ )
97
+
98
+ def fit(
99
+ self,
100
+ data: Any,
101
+ epochs: int = 20,
102
+ lr: float = 0.01,
103
+ weight_decay: float = 5e-4,
104
+ mask: Optional[str] = "train_mask",
105
+ val_mask: Optional[str] = "val_mask",
106
+ verbose: int = 1,
107
+ callbacks: Optional[List[Any]] = None,
108
+ ):
109
+ r"""Trains the node classifier on the provided graph data."""
110
+ if self.model is None:
111
+ self._init_model(data)
112
+
113
+ if not self._is_compiled:
114
+ opt = keras.optimizers.Adam(learning_rate=lr, weight_decay=weight_decay)
115
+ if self.multi_label:
116
+ loss = keras.losses.BinaryCrossentropy(from_logits=True)
117
+ metrics = [keras.metrics.BinaryAccuracy(name="acc")]
118
+ else:
119
+ loss = keras.losses.SparseCategoricalCrossentropy(from_logits=True)
120
+ metrics = [keras.metrics.SparseCategoricalAccuracy(name="acc")]
121
+
122
+ self.model.compile(
123
+ optimizer=opt,
124
+ loss=loss,
125
+ weighted_metrics=metrics,
126
+ )
127
+ self._is_compiled = True
128
+
129
+ # Generate training batches
130
+ if hasattr(data, "to_generator"):
131
+ gen = data.to_generator(mask=mask)
132
+ else:
133
+ inputs = self._extract_inputs(data)
134
+ y = ops.convert_to_tensor(data.y)
135
+ m = getattr(data, mask) if mask and hasattr(data, mask) else None
136
+ sample_weight = ops.cast(m, "float32") if m is not None else None
137
+ if sample_weight is not None: # the loss is the mean over the masked nodes, as in PyG
138
+ sample_weight = sample_weight * (ops.cast(ops.size(sample_weight), "float32") / ops.maximum(ops.sum(sample_weight), 1.0))
139
+
140
+ def gen_fn():
141
+ while True:
142
+ if sample_weight is not None:
143
+ yield inputs, y, sample_weight
144
+ else:
145
+ yield inputs, y
146
+
147
+ gen = gen_fn()
148
+
149
+ return self.model.fit(
150
+ gen,
151
+ steps_per_epoch=1,
152
+ epochs=epochs,
153
+ shuffle=False, # one full-graph batch per step
154
+ verbose=verbose,
155
+ callbacks=callbacks,
156
+ )
157
+
158
+ def predict_proba(self, data: Any, mask: Optional[str] = None):
159
+ r"""Returns class probabilities for nodes."""
160
+ if self.model is None:
161
+ raise RuntimeError("Model is not initialized. Fit or load a model first.")
162
+
163
+ inputs = self._extract_inputs(data)
164
+ logits = self.model(inputs, training=False)
165
+
166
+ if self.multi_label:
167
+ probs = ops.sigmoid(logits)
168
+ else:
169
+ probs = ops.softmax(logits, axis=-1)
170
+
171
+ if mask is not None and hasattr(data, mask):
172
+ m = getattr(data, mask)
173
+ probs = probs[m]
174
+ return probs
175
+
176
+ def predict(self, data: Any, mask: Optional[str] = None):
177
+ r"""Predicts discrete class labels for nodes."""
178
+ probs = self.predict_proba(data, mask=mask)
179
+ if self.multi_label:
180
+ return ops.cast(probs > 0.5, "int64")
181
+ return ops.argmax(probs, axis=-1)
182
+
183
+ def evaluate(self, data: Any, mask: Optional[str] = "test_mask") -> Dict[str, float]:
184
+ r"""Evaluates classification accuracy on a given mask."""
185
+ pred = self.predict(data, mask=mask)
186
+ y = data.y
187
+ if mask is not None and hasattr(data, mask):
188
+ m = getattr(data, mask)
189
+ y = y[m]
190
+
191
+ y_cast = ops.cast(y, "int64")
192
+ pred_cast = ops.cast(pred, "int64")
193
+ acc = float(ops.convert_to_numpy(ops.mean(ops.cast(pred_cast == y_cast, "float32"))))
194
+ return {"accuracy": acc}
@@ -0,0 +1,138 @@
1
+ """High-level Node Regression Task."""
2
+
3
+ from typing import Any, Dict, List, Optional, Union
4
+ import keras
5
+ from keras import ops
6
+
7
+ from k3_node.tasks.base import BaseTask
8
+ from k3_node.tasks.backbone_resolver import resolve_backbone
9
+
10
+
11
+ class NodeRegressor(BaseTask):
12
+ r"""High-level estimator for node regression tasks.
13
+
14
+ Args:
15
+ backbone: Model architecture string (``"gcn"``, ``"gat"``, ``"sage"``,
16
+ ``"gin"``, ``"pna"``, ``"mlp"``, etc.) or a custom :class:`keras.Model`.
17
+ (default: ``"gcn"``)
18
+ in_channels (int, optional): Size of input node features.
19
+ hidden_channels (int, optional): Dimensionality of hidden node features. (default: ``64``)
20
+ out_channels (int, optional): Number of continuous target variables. (default: ``1``)
21
+ num_layers (int, optional): Number of message passing layers. (default: ``2``)
22
+ dropout (float, optional): Dropout probability. (default: ``0.0``)
23
+ loss: Regression loss (``"mse"``, ``"mae"``, or a Keras loss instance). (default: ``"mse"``)
24
+ **backbone_kwargs: Additional arguments forwarded to the backbone constructor.
25
+ """
26
+
27
+ def __init__(
28
+ self,
29
+ backbone: Union[str, keras.Model] = "gcn",
30
+ in_channels: Optional[int] = None,
31
+ hidden_channels: int = 64,
32
+ out_channels: int = 1,
33
+ num_layers: int = 2,
34
+ dropout: float = 0.0,
35
+ loss: str = "mse",
36
+ **backbone_kwargs,
37
+ ):
38
+ super().__init__()
39
+ self.backbone = backbone
40
+ self.in_channels = in_channels
41
+ self.hidden_channels = hidden_channels
42
+ self.out_channels = out_channels
43
+ self.num_layers = num_layers
44
+ self.dropout = dropout
45
+ self.loss_name = loss
46
+ self.backbone_kwargs = backbone_kwargs
47
+
48
+ if isinstance(backbone, keras.Model):
49
+ self.model = backbone
50
+
51
+ def _init_model(self, data: Any):
52
+ in_c = self.in_channels
53
+ if in_c is None:
54
+ if hasattr(data, "num_node_features") and data.num_node_features > 0:
55
+ in_c = data.num_node_features
56
+ elif hasattr(data, "num_features") and data.num_features > 0:
57
+ in_c = data.num_features
58
+ elif hasattr(data, "x") and data.x is not None:
59
+ in_c = int(ops.shape(data.x)[-1])
60
+ else:
61
+ raise ValueError("Could not automatically infer in_channels from data.")
62
+
63
+ self.in_channels = in_c
64
+ self.model = resolve_backbone(
65
+ self.backbone,
66
+ in_channels=in_c,
67
+ out_channels=self.out_channels,
68
+ hidden_channels=self.hidden_channels,
69
+ num_layers=self.num_layers,
70
+ dropout=self.dropout,
71
+ **self.backbone_kwargs,
72
+ )
73
+
74
+ def fit(
75
+ self,
76
+ data: Any,
77
+ epochs: int = 20,
78
+ lr: float = 0.01,
79
+ mask: Optional[str] = "train_mask",
80
+ verbose: int = 1,
81
+ callbacks: Optional[List[Any]] = None,
82
+ ):
83
+ r"""Trains the node regressor on the provided graph data."""
84
+ if self.model is None:
85
+ self._init_model(data)
86
+
87
+ if not self._is_compiled:
88
+ loss = keras.losses.MeanSquaredError() if self.loss_name == "mse" else keras.losses.MeanAbsoluteError()
89
+ self.model.compile(
90
+ optimizer=keras.optimizers.Adam(learning_rate=lr),
91
+ loss=loss,
92
+ weighted_metrics=[keras.metrics.MeanAbsoluteError(name="mae")],
93
+ )
94
+ self._is_compiled = True
95
+
96
+ inputs = self._extract_inputs(data)
97
+ y = ops.cast(data.y, "float32")
98
+ m = getattr(data, mask) if mask and hasattr(data, mask) else None
99
+ sample_weight = ops.cast(m, "float32") if m is not None else None
100
+ if sample_weight is not None: # the loss is the mean over the masked nodes, as in PyG
101
+ sample_weight = sample_weight * (ops.cast(ops.size(sample_weight), "float32") / ops.maximum(ops.sum(sample_weight), 1.0))
102
+
103
+ def gen_fn():
104
+ while True:
105
+ if sample_weight is not None:
106
+ yield inputs, y, sample_weight
107
+ else:
108
+ yield inputs, y
109
+
110
+ return self.model.fit(
111
+ gen_fn(),
112
+ steps_per_epoch=1,
113
+ epochs=epochs,
114
+ shuffle=False, # one full-graph batch per step
115
+ verbose=verbose,
116
+ callbacks=callbacks,
117
+ )
118
+
119
+ def predict(self, data: Any, mask: Optional[str] = None):
120
+ r"""Returns continuous predictions for nodes."""
121
+ if self.model is None:
122
+ raise RuntimeError("Model is not initialized.")
123
+ inputs = self._extract_inputs(data)
124
+ pred = self.model(inputs, training=False)
125
+ if mask is not None and hasattr(data, mask):
126
+ pred = pred[getattr(data, mask)]
127
+ return pred
128
+
129
+ def evaluate(self, data: Any, mask: Optional[str] = "test_mask") -> Dict[str, float]:
130
+ r"""Evaluates Mean Absolute Error and Mean Squared Error."""
131
+ pred = self.predict(data, mask=mask)
132
+ y = data.y
133
+ if mask is not None and hasattr(data, mask):
134
+ y = y[getattr(data, mask)]
135
+ diff = ops.cast(pred, "float32") - ops.cast(y, "float32")
136
+ mae = float(ops.convert_to_numpy(ops.mean(ops.abs(diff))))
137
+ mse = float(ops.convert_to_numpy(ops.mean(ops.square(diff))))
138
+ return {"mae": mae, "mse": mse, "loss": mse}
@@ -0,0 +1,319 @@
1
+ """Unit tests for K3-Node high-level Task APIs."""
2
+
3
+ import os
4
+ import tempfile
5
+ import pytest
6
+ import numpy as np
7
+ import keras
8
+ from keras import ops
9
+
10
+ import k3_node
11
+ from k3_node.data import Data
12
+ from k3_node.tasks import (
13
+ BaseTask,
14
+ resolve_backbone,
15
+ NodeClassifier,
16
+ NodeRegressor,
17
+ GraphClassifier,
18
+ GraphRegressor,
19
+ LinkPredictor,
20
+ )
21
+
22
+
23
+ def _create_synthetic_node_data(num_nodes=24, in_channels=8, num_classes=3, multi_label=False):
24
+ x = np.random.randn(num_nodes, in_channels).astype("float32")
25
+ edges_src = np.arange(num_nodes - 1, dtype="int64")
26
+ edges_dst = np.arange(1, num_nodes, dtype="int64")
27
+ edge_index = np.stack([edges_src, edges_dst], axis=0)
28
+
29
+ if multi_label:
30
+ y = np.random.randint(0, 2, size=(num_nodes, num_classes)).astype("float32")
31
+ else:
32
+ y = np.random.randint(0, num_classes, size=(num_nodes,)).astype("int64")
33
+
34
+ train_mask = np.zeros(num_nodes, dtype=bool)
35
+ train_mask[: num_nodes // 2] = True
36
+ val_mask = np.zeros(num_nodes, dtype=bool)
37
+ val_mask[num_nodes // 2 : 3 * num_nodes // 4] = True
38
+ test_mask = ~(train_mask | val_mask)
39
+
40
+ return Data(
41
+ x=x,
42
+ edge_index=edge_index,
43
+ y=y,
44
+ train_mask=train_mask,
45
+ val_mask=val_mask,
46
+ test_mask=test_mask,
47
+ )
48
+
49
+
50
+ def _create_synthetic_graph_dataset(num_graphs=12, nodes_per_graph=6, in_channels=8, is_regression=False):
51
+ dataset = []
52
+ for i in range(num_graphs):
53
+ x = np.random.randn(nodes_per_graph, in_channels).astype("float32")
54
+ edges_src = np.arange(nodes_per_graph - 1, dtype="int64")
55
+ edges_dst = np.arange(1, nodes_per_graph, dtype="int64")
56
+ edge_index = np.stack([edges_src, edges_dst], axis=0)
57
+
58
+ if is_regression:
59
+ y = np.array([float(np.mean(x))], dtype="float32")
60
+ else:
61
+ y = np.array(i % 2, dtype="int64")
62
+
63
+ dataset.append(Data(x=x, edge_index=edge_index, y=y))
64
+ return dataset
65
+
66
+
67
+ # ==============================================================================
68
+ # 1. Top-Level Imports & Base Task Tests
69
+ # ==============================================================================
70
+
71
+ def test_tasks_module_exports():
72
+ """Verify tasks are exported at both k3_node.tasks and k3_node top level."""
73
+ assert hasattr(k3_node, "tasks")
74
+ assert hasattr(k3_node, "NodeClassifier")
75
+ assert hasattr(k3_node, "NodeRegressor")
76
+ assert hasattr(k3_node, "GraphClassifier")
77
+ assert hasattr(k3_node, "GraphRegressor")
78
+ assert hasattr(k3_node, "LinkPredictor")
79
+
80
+ assert k3_node.NodeClassifier is NodeClassifier
81
+ assert k3_node.NodeRegressor is NodeRegressor
82
+ assert k3_node.GraphClassifier is GraphClassifier
83
+ assert k3_node.GraphRegressor is GraphRegressor
84
+ assert k3_node.LinkPredictor is LinkPredictor
85
+
86
+
87
+ def test_base_task():
88
+ """Test BaseTask compile, extract_inputs, save and load."""
89
+ dummy_model = keras.Sequential([keras.layers.Dense(4)])
90
+ task = BaseTask(model=dummy_model)
91
+ task.compile(optimizer="adam", loss="mse")
92
+ assert task._is_compiled
93
+
94
+ # extract_inputs
95
+ data = _create_synthetic_node_data(num_nodes=5, in_channels=4)
96
+ extracted = task._extract_inputs(data)
97
+ assert isinstance(extracted, tuple)
98
+
99
+ # save and load
100
+ with tempfile.TemporaryDirectory() as tmpdir:
101
+ filepath = os.path.join(tmpdir, "test_task.keras")
102
+ task.save(filepath)
103
+ loaded = BaseTask.load(filepath)
104
+ assert loaded.model is not None
105
+
106
+
107
+ def test_resolve_backbone():
108
+ """Test backbone resolver with valid strings, custom models, and errors."""
109
+ for name in ["gcn", "gat", "sage", "gin", "mlp"]:
110
+ model = resolve_backbone(name, in_channels=8, hidden_channels=16, out_channels=4, num_layers=2)
111
+ assert isinstance(model, keras.Model)
112
+
113
+ # Custom model pass-through
114
+ custom_model = keras.Sequential([keras.layers.Dense(4)])
115
+ resolved_custom = resolve_backbone(custom_model, in_channels=8, out_channels=4)
116
+ assert resolved_custom is custom_model
117
+
118
+ # Unknown backbone
119
+ with pytest.raises(ValueError):
120
+ resolve_backbone("non_existent_backbone", in_channels=8, out_channels=4)
121
+
122
+ # Invalid type
123
+ with pytest.raises(TypeError):
124
+ resolve_backbone(12345, in_channels=8, out_channels=4)
125
+
126
+
127
+ # ==============================================================================
128
+ # 2. NodeClassifier Tests
129
+ # ==============================================================================
130
+
131
+ @pytest.mark.parametrize("backbone", ["gcn", "sage"])
132
+ def test_node_classifier_fit_predict_evaluate(backbone):
133
+ data = _create_synthetic_node_data(num_nodes=20, in_channels=8, num_classes=3)
134
+ clf = NodeClassifier(backbone=backbone, hidden_channels=16, num_layers=2, dropout=0.0)
135
+
136
+ # Fit
137
+ clf.fit(data, epochs=2, lr=0.01, verbose=0)
138
+ assert clf.model is not None
139
+
140
+ # Predict discrete labels
141
+ preds = clf.predict(data, mask="test_mask")
142
+ assert ops.shape(preds)[0] == int(ops.convert_to_numpy(ops.sum(ops.cast(data.test_mask, "int32"))))
143
+
144
+ # Predict probabilities
145
+ probs = clf.predict_proba(data, mask="test_mask")
146
+ assert ops.shape(probs)[-1] == 3
147
+
148
+ # Evaluate
149
+ metrics = clf.evaluate(data, mask="test_mask")
150
+ assert "accuracy" in metrics
151
+ assert 0.0 <= metrics["accuracy"] <= 1.0
152
+
153
+
154
+ def test_node_classifier_multi_label():
155
+ data = _create_synthetic_node_data(num_nodes=20, in_channels=8, num_classes=4, multi_label=True)
156
+ clf = NodeClassifier(backbone="gcn", hidden_channels=16, num_layers=2, multi_label=True)
157
+ clf.fit(data, epochs=2, verbose=0)
158
+
159
+ probs = clf.predict_proba(data, mask="test_mask")
160
+ assert ops.shape(probs)[-1] == 4
161
+
162
+ preds = clf.predict(data, mask="test_mask")
163
+ assert ops.shape(preds)[-1] == 4
164
+
165
+
166
+ # ==============================================================================
167
+ # 3. NodeRegressor Tests
168
+ # ==============================================================================
169
+
170
+ def test_node_regressor_fit_predict_evaluate():
171
+ data = _create_synthetic_node_data(num_nodes=20, in_channels=8, num_classes=1)
172
+ data.y = np.random.randn(20, 1).astype("float32")
173
+
174
+ reg = NodeRegressor(backbone="gat", hidden_channels=16, num_layers=2)
175
+ reg.fit(data, epochs=2, lr=0.01, verbose=0)
176
+
177
+ preds = reg.predict(data, mask="test_mask")
178
+ assert ops.shape(preds)[0] == int(ops.convert_to_numpy(ops.sum(ops.cast(data.test_mask, "int32"))))
179
+
180
+ metrics = reg.evaluate(data, mask="test_mask")
181
+ assert "mae" in metrics
182
+ assert "mse" in metrics
183
+ assert "loss" in metrics
184
+ assert metrics["mae"] >= 0.0
185
+
186
+
187
+ # ==============================================================================
188
+ # 4. GraphClassifier Tests
189
+ # ==============================================================================
190
+
191
+ @pytest.mark.parametrize("pooling", ["mean", "add", "max"])
192
+ def test_graph_classifier_fit_predict_evaluate(pooling):
193
+ dataset = _create_synthetic_graph_dataset(num_graphs=10, nodes_per_graph=5, in_channels=6)
194
+ clf = GraphClassifier(
195
+ backbone="gin",
196
+ hidden_channels=16,
197
+ num_layers=2,
198
+ pooling=pooling,
199
+ dropout=0.0,
200
+ )
201
+
202
+ clf.fit(dataset, epochs=2, batch_size=4, verbose=0)
203
+
204
+ # Predict
205
+ preds = clf.predict(dataset[:4], batch_size=4)
206
+ assert ops.shape(preds)[0] == 4
207
+
208
+ # Predict proba
209
+ probs = clf.predict_proba(dataset[:4], batch_size=4)
210
+ assert ops.shape(probs)[0] == 4
211
+ assert ops.shape(probs)[1] == 2
212
+
213
+ # Evaluate
214
+ metrics = clf.evaluate(dataset[:4], batch_size=4)
215
+ assert "accuracy" in metrics
216
+ assert 0.0 <= metrics["accuracy"] <= 1.0
217
+
218
+
219
+ # ==============================================================================
220
+ # 5. GraphRegressor Tests
221
+ # ==============================================================================
222
+
223
+ @pytest.mark.parametrize("loss_name", ["mae", "mse"])
224
+ def test_graph_regressor_fit_predict_evaluate(loss_name):
225
+ dataset = _create_synthetic_graph_dataset(num_graphs=10, nodes_per_graph=5, in_channels=6, is_regression=True)
226
+ reg = GraphRegressor(
227
+ backbone="sage",
228
+ hidden_channels=16,
229
+ num_layers=2,
230
+ pooling="mean",
231
+ loss=loss_name,
232
+ )
233
+
234
+ reg.fit(dataset, epochs=2, batch_size=4, verbose=0)
235
+
236
+ preds = reg.predict(dataset[:4], batch_size=4)
237
+ assert ops.shape(preds)[0] == 4
238
+
239
+ metrics = reg.evaluate(dataset[:4], batch_size=4)
240
+ assert "mae" in metrics
241
+ assert "mse" in metrics
242
+ assert "loss" in metrics
243
+ assert metrics["mae"] >= 0.0
244
+
245
+
246
+ # ==============================================================================
247
+ # 6. LinkPredictor Tests
248
+ # ==============================================================================
249
+
250
+ @pytest.mark.parametrize("decoder", ["inner_product", "cosine", "mlp"])
251
+ def test_link_predictor_decoders_and_fit(decoder):
252
+ data = _create_synthetic_node_data(num_nodes=16, in_channels=8)
253
+ lp = LinkPredictor(
254
+ backbone="gcn",
255
+ hidden_channels=16,
256
+ out_channels=16,
257
+ num_layers=2,
258
+ decoder=decoder,
259
+ )
260
+
261
+ # Dynamic negative sampling training
262
+ lp.fit(data, epochs=2, lr=0.01, neg_ratio=1.0, verbose=0)
263
+
264
+ # Encode node embeddings
265
+ z = lp.encode(data)
266
+ assert ops.shape(z) == (16, 16)
267
+
268
+ # Predict proba & binary labels
269
+ query_edges = np.array([[0, 1, 2], [1, 2, 3]], dtype="int32")
270
+ probs = lp.predict_proba(data, edge_label_index=query_edges)
271
+ assert ops.shape(probs) == (3,)
272
+
273
+ preds = lp.predict(data, edge_label_index=query_edges, threshold=0.5)
274
+ assert ops.shape(preds) == (3,)
275
+
276
+ # Evaluate
277
+ metrics = lp.evaluate(data, edge_label_index=query_edges, edge_label=[1, 1, 0])
278
+ assert "accuracy" in metrics
279
+ assert 0.0 <= metrics["accuracy"] <= 1.0
280
+
281
+
282
+ def test_link_predictor_explicit_labels():
283
+ data = _create_synthetic_node_data(num_nodes=16, in_channels=8)
284
+ edge_label_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int32")
285
+ edge_label = np.array([1, 1, 0, 0], dtype="float32")
286
+
287
+ lp = LinkPredictor(backbone="gcn", hidden_channels=16, out_channels=16)
288
+ lp.fit(data, edge_label_index=edge_label_index, edge_label=edge_label, epochs=2, verbose=0)
289
+
290
+ metrics = lp.evaluate(data, edge_label_index=edge_label_index, edge_label=edge_label)
291
+ assert "accuracy" in metrics
292
+ assert "auc" in metrics
293
+ assert "ap" in metrics
294
+
295
+
296
+ # ==============================================================================
297
+ # 7. Applications API Tests
298
+ # ==============================================================================
299
+
300
+ def test_applications_api_exports():
301
+ """Verify k3_node.applications submodules and model exports."""
302
+ import k3_node.applications as apps
303
+ assert hasattr(apps, "chemistry")
304
+ assert hasattr(apps, "materials")
305
+ assert hasattr(apps, "bio")
306
+
307
+ # Chemistry exports
308
+ assert hasattr(apps.chemistry, "AttentiveFP")
309
+ assert hasattr(apps.chemistry, "SchNet")
310
+ assert hasattr(apps.chemistry, "DimeNetPlusPlus")
311
+
312
+ # Materials exports
313
+ assert hasattr(apps.materials, "MEGNet")
314
+ assert hasattr(apps.materials, "M3GNet")
315
+ assert hasattr(apps.materials, "CHGNet")
316
+
317
+ # Bio exports
318
+ assert hasattr(apps.bio, "UniMolDockingModel")
319
+ assert hasattr(apps.bio, "DockingPoseModelV2")