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,133 @@
1
+ """Model card generator for Hugging Face Hub integration."""
2
+
3
+ from typing import Any, Dict, Optional
4
+
5
+
6
+ def generate_model_card(
7
+ task_type: str,
8
+ backbone: str,
9
+ config: Dict[str, Any],
10
+ metrics: Optional[Dict[str, float]] = None,
11
+ dataset_name: Optional[str] = None,
12
+ repo_id: Optional[str] = None,
13
+ license: str = "mit",
14
+ ) -> str:
15
+ r"""Generates a standard Hugging Face Model Card with YAML frontmatter
16
+ and markdown documentation for a K3-Node GNN model.
17
+
18
+ Args:
19
+ task_type: Name of the task or model (e.g. 'NodeClassifier', 'SchNet').
20
+ backbone: Name of the backbone architecture (e.g. 'gcn', 'schnet', 'chgnet').
21
+ config: Dictionary containing model architecture and training hyperparameters.
22
+ metrics: Optional dictionary of evaluation metrics (e.g. {'accuracy': 0.82}).
23
+ dataset_name: Optional name of the dataset the model was trained on.
24
+ repo_id: Optional repository ID on Hugging Face Hub.
25
+ license: Open-source license tag. (default: 'mit')
26
+
27
+ Returns:
28
+ Formatted markdown string representing the README.md model card.
29
+ """
30
+ model_class = config.get("model_class", task_type)
31
+ task_slug = task_type.lower().replace("classifier", "-classification").replace("regressor", "-regression")
32
+ title = repo_id if repo_id else f"{backbone.upper()} {task_type}"
33
+ dataset_str = f" on `{dataset_name}`" if dataset_name else ""
34
+
35
+ # YAML Frontmatter
36
+ frontmatter = f"""---
37
+ language:
38
+ - en
39
+ license: {license}
40
+ library_name: keras
41
+ tags:
42
+ - graph-machine-learning
43
+ - gnn
44
+ - k3-node
45
+ - keras-3
46
+ - multi-backend
47
+ - {task_slug}
48
+ pipeline_tag: graph-ml
49
+ ---
50
+ """
51
+
52
+ # Model Details
53
+ content = f"""# {title}
54
+
55
+ This is a **{task_type}** Graph Neural Network model{dataset_str} built with [**K3-Node**](https://github.com/anas-rz/k3-node) and **Keras 3**.
56
+
57
+ It runs natively and seamlessly across **PyTorch**, **JAX**, and **TensorFlow** backends.
58
+
59
+ ## Model Details
60
+
61
+ - **Model / Task**: `{task_type}`
62
+ - **Architecture**: `{backbone}`
63
+ - **Library**: `k3-node` (Keras 3)
64
+ - **Input Channels**: `{config.get('in_channels', 'Auto')}`
65
+ - **Hidden Channels**: `{config.get('hidden_channels', 64)}`
66
+ - **Output / Classes**: `{config.get('num_classes') or config.get('out_channels', 'Auto')}`
67
+ - **Number of Layers**: `{config.get('num_layers', 2)}`
68
+ - **Dropout**: `{config.get('dropout', 0.0)}`
69
+ """
70
+
71
+ # Optional metrics
72
+ if metrics:
73
+ content += "\n## Evaluation Metrics\n\n| Metric | Value |\n| :--- | :--- |\n"
74
+ for k, v in metrics.items():
75
+ if isinstance(v, float):
76
+ content += f"| {k} | {v:.4f} |\n"
77
+ else:
78
+ content += f"| {k} | {v} |\n"
79
+
80
+ # Usage code snippet
81
+ target_repo = repo_id or "username/model-repo"
82
+ is_task = task_type in ("NodeClassifier", "GraphClassifier", "GraphRegressor", "LinkPredictor")
83
+
84
+ if is_task:
85
+ usage_code = f"""from k3_node.tasks import {task_type}
86
+
87
+ # Load the pretrained model directly from Hugging Face Hub
88
+ model = {task_type}.from_pretrained("{target_repo}")
89
+
90
+ # Run predictions on your graph data
91
+ predictions = model.predict(data)"""
92
+ else:
93
+ usage_code = f"""import k3_node as k3
94
+
95
+ # Load pre-trained weights with one line
96
+ model = k3.models.{model_class}.from_pretrained("{target_repo}")
97
+
98
+ # Run predictions directly on graph or molecular data
99
+ predictions = model.predict(data)"""
100
+
101
+ content += f"""
102
+ ## Usage
103
+
104
+ Install `k3-node` with your preferred backend (PyTorch, JAX, or TensorFlow):
105
+
106
+ ```bash
107
+ pip install k3-node huggingface_hub
108
+ ```
109
+
110
+ ### Loading & Inference
111
+
112
+ ```python
113
+ import os
114
+ os.environ["KERAS_BACKEND"] = "torch" # or "jax", "tensorflow"
115
+
116
+ {usage_code}
117
+ ```
118
+
119
+ ## Framework & Citation
120
+
121
+ This model was trained with **K3-Node**, the multi-backend Graph Neural Network framework built on Keras 3.
122
+
123
+ ```bibtex
124
+ @software{{k3_node,
125
+ author = {{Muhammad Anas Raza}},
126
+ title = {{K3-Node: Multi-Backend Graph Neural Networks on Keras 3}},
127
+ year = {{2026}},
128
+ url = {{https://github.com/anas-rz/k3-node}}
129
+ }}
130
+ ```
131
+ """
132
+
133
+ return frontmatter + "\n" + content.strip() + "\n"
@@ -0,0 +1,419 @@
1
+ """Tests for Hugging Face Hub integration in K3-Node."""
2
+
3
+ import os
4
+ import tempfile
5
+ from pathlib import Path
6
+ from unittest.mock import MagicMock, patch
7
+
8
+ import pytest
9
+ import numpy as np
10
+ from keras import ops
11
+
12
+ import k3_node
13
+ from k3_node.data import Data
14
+ from k3_node.tasks import (
15
+ NodeClassifier,
16
+ GraphClassifier,
17
+ GraphRegressor,
18
+ LinkPredictor,
19
+ )
20
+ from k3_node.hub import (
21
+ K3NodeHubMixin,
22
+ from_pretrained,
23
+ push_to_hub,
24
+ save_pretrained,
25
+ generate_model_card,
26
+ generate_dataset_card,
27
+ save_graph_dataset,
28
+ load_graph_dataset,
29
+ push_dataset_to_hub,
30
+ load_dataset_from_hub,
31
+ )
32
+
33
+
34
+ def _create_synthetic_node_data(num_nodes=12, in_channels=8, num_classes=2):
35
+ x = ops.convert_to_tensor(np.random.randn(num_nodes, in_channels), dtype="float32")
36
+ src = np.arange(num_nodes, dtype="int64")
37
+ dst = (src + 1) % num_nodes
38
+ edge_index = ops.convert_to_tensor(np.stack([src, dst], axis=0), dtype="int64")
39
+ y = ops.convert_to_tensor(np.random.randint(0, num_classes, size=(num_nodes,)), dtype="int64")
40
+ return Data(x=x, edge_index=edge_index, y=y)
41
+
42
+
43
+ def _create_synthetic_graph_dataset(num_graphs=6, nodes_per_graph=4, in_channels=6):
44
+ graphs = []
45
+ for i in range(num_graphs):
46
+ x = ops.convert_to_tensor(np.random.randn(nodes_per_graph, in_channels), dtype="float32")
47
+ edge_index = ops.convert_to_tensor(np.array([[0, 1, 2], [1, 2, 3]], dtype="int64"), dtype="int64")
48
+ y = ops.convert_to_tensor(np.array([i % 2], dtype="int64"), dtype="int64")
49
+ graphs.append(Data(x=x, edge_index=edge_index, y=y))
50
+ return graphs
51
+
52
+
53
+ def test_hub_module_exports():
54
+ assert hasattr(k3_node, "hub")
55
+ assert hasattr(k3_node, "from_pretrained")
56
+ assert hasattr(k3_node, "push_to_hub")
57
+ assert hasattr(k3_node, "save_pretrained")
58
+ assert hasattr(k3_node, "load_dataset_from_hub")
59
+ assert hasattr(k3_node, "push_dataset_to_hub")
60
+
61
+
62
+ def test_generate_model_card():
63
+ config = {
64
+ "task_type": "NodeClassifier",
65
+ "backbone": "gcn",
66
+ "in_channels": 16,
67
+ "hidden_channels": 32,
68
+ "out_channels": 3,
69
+ "num_layers": 2,
70
+ "dropout": 0.2,
71
+ }
72
+ card = generate_model_card(
73
+ task_type="NodeClassifier",
74
+ backbone="gcn",
75
+ config=config,
76
+ metrics={"accuracy": 0.845, "loss": 0.32},
77
+ dataset_name="Cora",
78
+ repo_id="anas-rz/cora-gcn",
79
+ )
80
+ assert "pipeline_tag: graph-ml" in card
81
+ assert "graph-machine-learning" in card
82
+ assert "# anas-rz/cora-gcn" in card
83
+ assert "Cora" in card
84
+ assert "0.8450" in card
85
+ assert "NodeClassifier.from_pretrained" in card
86
+
87
+
88
+ def test_generate_dataset_card():
89
+ data = _create_synthetic_node_data()
90
+ card = generate_dataset_card(data, repo_id="anas-rz/synthetic-graph", description="Synthetic graph test dataset")
91
+ assert "pipeline_tag: graph-ml" in card or "graph-dataset" in card
92
+ assert "Total Graphs" in card
93
+ assert "anas-rz/synthetic-graph" in card
94
+
95
+
96
+ def test_save_and_from_pretrained_node_classifier():
97
+ data = _create_synthetic_node_data(num_nodes=10, in_channels=8, num_classes=2)
98
+ clf = NodeClassifier(backbone="gcn", in_channels=8, out_channels=2, hidden_channels=16, num_layers=2)
99
+ clf.fit(data, epochs=2, verbose=0)
100
+ orig_preds = clf.predict(data)
101
+
102
+ with tempfile.TemporaryDirectory() as tmpdir:
103
+ clf.save_pretrained(tmpdir, metrics={"accuracy": 1.0})
104
+ p = Path(tmpdir)
105
+ assert (p / "config.json").exists()
106
+ assert (p / "model.weights.h5").exists()
107
+ assert (p / "README.md").exists()
108
+
109
+ # Load with specific class
110
+ loaded_clf = NodeClassifier.from_pretrained(tmpdir)
111
+ assert isinstance(loaded_clf, NodeClassifier)
112
+ new_preds = loaded_clf.predict(data)
113
+ np.testing.assert_array_equal(ops.convert_to_numpy(orig_preds), ops.convert_to_numpy(new_preds))
114
+
115
+ # Load with generic from_pretrained
116
+ generic_loaded = from_pretrained(tmpdir)
117
+ assert isinstance(generic_loaded, NodeClassifier)
118
+ gen_preds = generic_loaded.predict(data)
119
+ np.testing.assert_array_equal(ops.convert_to_numpy(orig_preds), ops.convert_to_numpy(gen_preds))
120
+
121
+
122
+ def test_save_and_from_pretrained_graph_classifier():
123
+ dataset = _create_synthetic_graph_dataset(num_graphs=4, nodes_per_graph=4, in_channels=6)
124
+ clf = GraphClassifier(backbone="gin", in_channels=6, num_classes=2, hidden_channels=16, num_layers=2)
125
+ clf.fit(dataset, epochs=2, batch_size=2, verbose=0)
126
+ orig_preds = clf.predict(dataset, batch_size=2)
127
+
128
+ with tempfile.TemporaryDirectory() as tmpdir:
129
+ clf.save_pretrained(tmpdir)
130
+ loaded = GraphClassifier.from_pretrained(tmpdir)
131
+ assert isinstance(loaded, GraphClassifier)
132
+ new_preds = loaded.predict(dataset, batch_size=2)
133
+ np.testing.assert_array_equal(ops.convert_to_numpy(orig_preds), ops.convert_to_numpy(new_preds))
134
+
135
+
136
+ def test_save_and_from_pretrained_graph_regressor():
137
+ dataset = _create_synthetic_graph_dataset(num_graphs=4, nodes_per_graph=4, in_channels=6)
138
+ reg = GraphRegressor(backbone="gin", in_channels=6, out_channels=1, hidden_channels=16, num_layers=2)
139
+ reg.fit(dataset, epochs=2, batch_size=2, verbose=0)
140
+ orig_preds = reg.predict(dataset, batch_size=2)
141
+
142
+ with tempfile.TemporaryDirectory() as tmpdir:
143
+ reg.save_pretrained(tmpdir)
144
+ loaded = GraphRegressor.from_pretrained(tmpdir)
145
+ assert isinstance(loaded, GraphRegressor)
146
+ new_preds = loaded.predict(dataset, batch_size=2)
147
+ np.testing.assert_allclose(ops.convert_to_numpy(orig_preds), ops.convert_to_numpy(new_preds), rtol=1e-5)
148
+
149
+
150
+ def test_save_and_from_pretrained_link_predictor():
151
+ data = _create_synthetic_node_data(num_nodes=8, in_channels=6)
152
+ lp = LinkPredictor(backbone="gcn", in_channels=6, hidden_channels=16, out_channels=16, num_layers=2)
153
+ lp.fit(data, epochs=2, verbose=0)
154
+ orig_probs = lp.predict_proba(data, edge_label_index=data.edge_index)
155
+
156
+ with tempfile.TemporaryDirectory() as tmpdir:
157
+ lp.save_pretrained(tmpdir)
158
+ loaded = LinkPredictor.from_pretrained(tmpdir)
159
+ assert isinstance(loaded, LinkPredictor)
160
+ new_probs = loaded.predict_proba(data, edge_label_index=data.edge_index)
161
+ np.testing.assert_allclose(ops.convert_to_numpy(orig_probs), ops.convert_to_numpy(new_probs), rtol=1e-4)
162
+
163
+
164
+ def test_save_and_load_single_graph_dataset():
165
+ data = _create_synthetic_node_data(num_nodes=8, in_channels=4)
166
+ with tempfile.TemporaryDirectory() as tmpdir:
167
+ path = os.path.join(tmpdir, "single_graph.npz")
168
+ saved_path = save_graph_dataset(data, path)
169
+ assert os.path.exists(saved_path)
170
+
171
+ loaded = load_graph_dataset(saved_path)
172
+ assert isinstance(loaded, Data)
173
+ np.testing.assert_allclose(ops.convert_to_numpy(data.x), ops.convert_to_numpy(loaded.x))
174
+ np.testing.assert_array_equal(ops.convert_to_numpy(data.edge_index), ops.convert_to_numpy(loaded.edge_index))
175
+ np.testing.assert_array_equal(ops.convert_to_numpy(data.y), ops.convert_to_numpy(loaded.y))
176
+
177
+
178
+ def test_save_and_load_graph_list_dataset():
179
+ graphs = _create_synthetic_graph_dataset(num_graphs=4, nodes_per_graph=3, in_channels=5)
180
+ with tempfile.TemporaryDirectory() as tmpdir:
181
+ path = os.path.join(tmpdir, "graph_list.npz")
182
+ saved_path = save_graph_dataset(graphs, path)
183
+ assert os.path.exists(saved_path)
184
+
185
+ loaded = load_graph_dataset(saved_path)
186
+ assert isinstance(loaded, list)
187
+ assert len(loaded) == 4
188
+ for orig, rec in zip(graphs, loaded):
189
+ np.testing.assert_allclose(ops.convert_to_numpy(orig.x), ops.convert_to_numpy(rec.x))
190
+ np.testing.assert_array_equal(ops.convert_to_numpy(orig.edge_index), ops.convert_to_numpy(rec.edge_index))
191
+
192
+
193
+ def test_push_to_hub_mocked():
194
+ data = _create_synthetic_node_data(num_nodes=6, in_channels=4)
195
+ clf = NodeClassifier(backbone="gcn", in_channels=4, out_channels=2, hidden_channels=8)
196
+ clf.fit(data, epochs=1, verbose=0)
197
+
198
+ with patch("huggingface_hub.HfApi") as MockApi:
199
+ mock_api_instance = MagicMock()
200
+ MockApi.return_value = mock_api_instance
201
+
202
+ url = clf.push_to_hub("test-user/test-gcn", token="dummy_token")
203
+ assert url == "https://huggingface.co/test-user/test-gcn"
204
+ mock_api_instance.create_repo.assert_called_once()
205
+ mock_api_instance.upload_folder.assert_called_once()
206
+ args, kwargs = mock_api_instance.upload_folder.call_args
207
+ assert kwargs.get("repo_id") == "test-user/test-gcn"
208
+ assert kwargs.get("repo_type") == "model"
209
+
210
+
211
+ def test_push_dataset_to_hub_mocked():
212
+ data = _create_synthetic_node_data(num_nodes=6, in_channels=4)
213
+
214
+ with patch("huggingface_hub.HfApi") as MockApi:
215
+ mock_api_instance = MagicMock()
216
+ MockApi.return_value = mock_api_instance
217
+
218
+ url = push_dataset_to_hub(data, "test-user/test-dataset", token="dummy_token")
219
+ assert url == "https://huggingface.co/datasets/test-user/test-dataset"
220
+ mock_api_instance.create_repo.assert_called_once()
221
+ mock_api_instance.upload_folder.assert_called_once()
222
+ args, kwargs = mock_api_instance.upload_folder.call_args
223
+ assert kwargs.get("repo_id") == "test-user/test-dataset"
224
+ assert kwargs.get("repo_type") == "dataset"
225
+
226
+
227
+ def test_load_from_hub_mocked():
228
+ data = _create_synthetic_node_data(num_nodes=8, in_channels=6, num_classes=2)
229
+ clf = NodeClassifier(backbone="gcn", in_channels=6, out_channels=2, hidden_channels=16)
230
+ clf.fit(data, epochs=1, verbose=0)
231
+
232
+ with tempfile.TemporaryDirectory() as tmpdir:
233
+ clf.save_pretrained(tmpdir)
234
+ config_file = os.path.join(tmpdir, "config.json")
235
+ weights_file = os.path.join(tmpdir, "model.weights.h5")
236
+
237
+ def mock_download(repo_id, filename, **kwargs):
238
+ if filename == "config.json":
239
+ return config_file
240
+ elif filename == "model.weights.h5":
241
+ return weights_file
242
+ raise FileNotFoundError(filename)
243
+
244
+ with patch("huggingface_hub.hf_hub_download", side_effect=mock_download):
245
+ loaded = NodeClassifier.from_pretrained("test-user/test-gcn")
246
+ assert isinstance(loaded, NodeClassifier)
247
+ preds = loaded.predict(data)
248
+ assert ops.shape(preds)[0] == 8
249
+
250
+
251
+ def test_load_dataset_from_hub_mocked():
252
+ data = _create_synthetic_node_data(num_nodes=6, in_channels=4)
253
+ with tempfile.TemporaryDirectory() as tmpdir:
254
+ dataset_path = os.path.join(tmpdir, "graph_data.npz")
255
+ save_graph_dataset(data, dataset_path)
256
+
257
+ with patch("huggingface_hub.hf_hub_download", return_value=dataset_path):
258
+ loaded = load_dataset_from_hub("test-user/test-graph-dataset")
259
+ assert isinstance(loaded, Data)
260
+ np.testing.assert_allclose(ops.convert_to_numpy(data.x), ops.convert_to_numpy(loaded.x))
261
+
262
+
263
+ def test_schnet_hub_save_load_predict():
264
+ """Test k3.models.SchNet.from_pretrained and model.predict(molecule_data)."""
265
+ from k3_node.models import SchNet
266
+
267
+ model = SchNet(
268
+ hidden_channels=16,
269
+ num_filters=16,
270
+ num_interactions=2,
271
+ num_gaussians=10,
272
+ cutoff=5.0,
273
+ )
274
+ z = ops.convert_to_tensor([1, 6, 8, 1])
275
+ pos = ops.convert_to_tensor([
276
+ [0.0, 0.0, 0.0],
277
+ [1.0, 0.0, 0.0],
278
+ [0.0, 1.0, 0.0],
279
+ [1.0, 1.0, 0.0],
280
+ ], dtype="float32")
281
+ molecule_data = Data(z=z, pos=pos)
282
+
283
+ energy = model.predict(molecule_data)
284
+ assert energy.shape == (1, 1)
285
+
286
+ with tempfile.TemporaryDirectory() as tmpdir:
287
+ model.save_pretrained(tmpdir, repo_id="k3-node/schnet-qm9")
288
+
289
+ # Load pre-trained weights with one line
290
+ reloaded = SchNet.from_pretrained(tmpdir)
291
+ assert isinstance(reloaded, SchNet)
292
+ energy_reloaded = reloaded.predict(molecule_data)
293
+ np.testing.assert_allclose(
294
+ ops.convert_to_numpy(energy),
295
+ ops.convert_to_numpy(energy_reloaded),
296
+ atol=1e-5,
297
+ )
298
+
299
+ # Test generic hub.from_pretrained
300
+ generic_loaded = from_pretrained(tmpdir)
301
+ assert isinstance(generic_loaded, SchNet)
302
+ np.testing.assert_allclose(
303
+ ops.convert_to_numpy(energy),
304
+ ops.convert_to_numpy(generic_loaded.predict(molecule_data)),
305
+ atol=1e-5,
306
+ )
307
+
308
+ # Push community checkpoints directly to the hub (mocked)
309
+ with patch("huggingface_hub.HfApi") as MockApi:
310
+ mock_api_instance = MagicMock()
311
+ MockApi.return_value = mock_api_instance
312
+ url = model.push_to_hub("k3-node/schnet-qm9", token="dummy_token")
313
+ assert url == "https://huggingface.co/k3-node/schnet-qm9"
314
+
315
+
316
+ def test_gcn_hub_save_load_predict():
317
+ """Test k3.models.GCN save_pretrained, from_pretrained, and predict."""
318
+ from k3_node.models import GCN
319
+
320
+ gcn = GCN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=3)
321
+ x = ops.zeros((4, 8), dtype="float32")
322
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
323
+ graph = Data(x=x, edge_index=edge_index)
324
+
325
+ p1 = gcn.predict(graph)
326
+ assert p1.shape == (4, 3)
327
+
328
+ with tempfile.TemporaryDirectory() as tmpdir:
329
+ gcn.save_pretrained(tmpdir)
330
+ reloaded = GCN.from_pretrained(tmpdir)
331
+ assert isinstance(reloaded, GCN)
332
+ p2 = reloaded.predict(graph)
333
+ np.testing.assert_allclose(
334
+ ops.convert_to_numpy(p1),
335
+ ops.convert_to_numpy(p2),
336
+ atol=1e-5,
337
+ )
338
+
339
+
340
+ def test_chgnet_hub_save_load_predict():
341
+ """Test k3.models.CHGNet.from_pretrained, predict, and push_to_hub."""
342
+ from k3_node.models.materials import CHGNet
343
+
344
+ model = CHGNet(
345
+ dim_atom_embedding=16,
346
+ dim_bond_embedding=16,
347
+ dim_angle_embedding=16,
348
+ num_blocks=2,
349
+ atom_conv_hidden_dims=(16,),
350
+ bond_conv_hidden_dims=(16,),
351
+ )
352
+
353
+ crystal = {
354
+ "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype=np.float32),
355
+ "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]], dtype=np.int32),
356
+ "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype=np.int32),
357
+ "node_type": np.array([6, 8, 1, 6], dtype=np.int32),
358
+ "batch": np.array([0, 0, 0, 0], dtype=np.int32),
359
+ "state_attr": np.array([[0.0, 0.0]], dtype=np.float32),
360
+ }
361
+
362
+ pred = model.predict(crystal)
363
+ assert pred.shape == () or pred.shape == (1,)
364
+
365
+ with tempfile.TemporaryDirectory() as tmpdir:
366
+ model.save_pretrained(tmpdir, repo_id="anas-rz/chgnet-mp-2026")
367
+ reloaded = CHGNet.from_pretrained(tmpdir)
368
+ assert isinstance(reloaded, CHGNet)
369
+ pred_reloaded = reloaded.predict(crystal)
370
+ np.testing.assert_allclose(
371
+ ops.convert_to_numpy(pred),
372
+ ops.convert_to_numpy(pred_reloaded),
373
+ atol=1e-5,
374
+ )
375
+
376
+ # Push community checkpoints directly to the hub (mocked)
377
+ with patch("huggingface_hub.HfApi") as MockApi:
378
+ mock_api_instance = MagicMock()
379
+ MockApi.return_value = mock_api_instance
380
+ url = model.push_to_hub("anas-rz/chgnet-mp-2026", token="dummy_token")
381
+ assert url == "https://huggingface.co/anas-rz/chgnet-mp-2026"
382
+
383
+
384
+
385
+ def test_hub_injection_keeps_model_specific_from_pretrained():
386
+ """Models that load original checkpoints must keep their own ``from_pretrained``."""
387
+ from k3_node.models import GraphMAE2, Graphormer, Graphormer3D
388
+
389
+ for cls in (GraphMAE2, Graphormer, Graphormer3D):
390
+ assert cls.from_pretrained.__func__ is vars(cls)["from_pretrained"].__func__, cls.__name__
391
+
392
+
393
+ def test_hub_injection_keeps_keras_batched_predict():
394
+ """Array inputs must still go through Keras' batched ``Model.predict``; graph inputs use the hub path."""
395
+ from k3_node.models import MLP, GCN
396
+
397
+ mlp = MLP([8, 16, 3])
398
+ x = np.random.randn(10, 8).astype("float32")
399
+ expected = ops.convert_to_numpy(mlp(x, training=False))
400
+ np.testing.assert_allclose(mlp.predict(x, batch_size=4, verbose=0), expected, rtol=1e-5, atol=1e-6)
401
+
402
+ gcn = GCN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=3)
403
+ graph = Data(x=x, edge_index=np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int32"))
404
+ assert tuple(ops.shape(gcn.predict(graph))) == (10, 3)
405
+
406
+
407
+ def test_backbone_kwargs_survive_save_and_load():
408
+ # `**backbone_kwargs` used to be passed back as a nested `backbone_kwargs=` argument on reload,
409
+ # silently dropping options such as the number of attention heads.
410
+ data = _create_synthetic_node_data(num_nodes=10, in_channels=8, num_classes=2)
411
+ clf = NodeClassifier(backbone="gat", in_channels=8, out_channels=2, hidden_channels=16, num_layers=2, heads=4)
412
+ clf.fit(data, epochs=1, verbose=0)
413
+ orig_preds = clf.predict(data)
414
+
415
+ with tempfile.TemporaryDirectory() as tmpdir:
416
+ clf.save_pretrained(tmpdir)
417
+ loaded = NodeClassifier.from_pretrained(tmpdir)
418
+ assert loaded.backbone_kwargs == {"heads": 4}
419
+ np.testing.assert_array_equal(ops.convert_to_numpy(orig_preds), ops.convert_to_numpy(loaded.predict(data)))
k3_node/io/__init__.py ADDED
@@ -0,0 +1,22 @@
1
+ from k3_node.io.fs import cp, exists, glob_files, isdir, isfile, rm
2
+ from k3_node.io.npz import parse_npz, read_npz
3
+ from k3_node.io.planetoid import read_planetoid_data
4
+ from k3_node.io.tu import read_tu_data
5
+ from k3_node.io.txt_array import parse_txt_array, read_txt_array
6
+
7
+ __all__ = [
8
+ "exists",
9
+ "isdir",
10
+ "isfile",
11
+ "cp",
12
+ "rm",
13
+ "glob_files",
14
+ "parse_txt_array",
15
+ "read_txt_array",
16
+ "parse_npz",
17
+ "read_npz",
18
+ "read_planetoid_data",
19
+ "read_tu_data",
20
+ ]
21
+
22
+ from k3_node.io.off import parse_off, read_off
k3_node/io/fs.py ADDED
@@ -0,0 +1,117 @@
1
+ import glob
2
+ import gzip
3
+ import os
4
+ import os.path as osp
5
+ import shutil
6
+ import ssl
7
+ import sys
8
+ import tarfile
9
+ import urllib.request
10
+ import zipfile
11
+ from typing import List, Optional
12
+
13
+
14
+ def exists(path: str) -> bool:
15
+ return osp.exists(path)
16
+
17
+
18
+ def isdir(path: str) -> bool:
19
+ return osp.isdir(path)
20
+
21
+
22
+ def isfile(path: str) -> bool:
23
+ return osp.isfile(path)
24
+
25
+
26
+ def is_url(url: str) -> bool:
27
+ return url.startswith("http://") or url.startswith("https://")
28
+
29
+
30
+ def download_url(
31
+ url: str,
32
+ folder: str,
33
+ filename: Optional[str] = None,
34
+ log: bool = True,
35
+ ) -> str:
36
+ r"""Downloads the content of an URL to a specific folder."""
37
+ if filename is None:
38
+ filename = url.rpartition("/")[2].split("?")[0]
39
+
40
+ os.makedirs(folder, exist_ok=True)
41
+ out_path = osp.join(folder, filename)
42
+
43
+ if osp.exists(out_path):
44
+ return out_path
45
+
46
+ if log and "PYTEST_CURRENT_TEST" not in os.environ:
47
+ print(f"Downloading {url}", file=sys.stderr)
48
+
49
+ req = urllib.request.Request(
50
+ url,
51
+ headers={"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64)"},
52
+ )
53
+ try:
54
+ context = ssl.create_default_context()
55
+ with urllib.request.urlopen(req, context=context) as response, open(out_path, "wb") as out_file:
56
+ shutil.copyfileobj(response, out_file)
57
+ except Exception:
58
+ # Fallback to unverified context if SSL verification fails
59
+ context = ssl._create_unverified_context()
60
+ with urllib.request.urlopen(req, context=context) as response, open(out_path, "wb") as out_file:
61
+ shutil.copyfileobj(response, out_file)
62
+
63
+ return out_path
64
+
65
+
66
+ def extract_archive(path: str, dst: str):
67
+ r"""Extracts an archive (zip, tar.gz, tgz, gz) to dst directory."""
68
+ os.makedirs(dst, exist_ok=True)
69
+ if path.endswith(".zip"):
70
+ with zipfile.ZipFile(path, "r") as zf:
71
+ zf.extractall(dst)
72
+ elif path.endswith(".tar.gz") or path.endswith(".tgz"):
73
+ with tarfile.open(path, "r:gz") as tf:
74
+ tf.extractall(dst)
75
+ elif path.endswith(".tar"):
76
+ with tarfile.open(path, "r:") as tf:
77
+ tf.extractall(dst)
78
+ elif path.endswith(".gz"):
79
+ out_name = osp.splitext(osp.basename(path))[0]
80
+ out_file = osp.join(dst, out_name)
81
+ with gzip.open(path, "rb") as f_in, open(out_file, "wb") as f_out:
82
+ shutil.copyfileobj(f_in, f_out)
83
+
84
+
85
+ def cp(src: str, dst: str, extract: bool = False, log: bool = True):
86
+ if is_url(src):
87
+ # Determine destination folder and filename
88
+ if dst.endswith("/") or osp.isdir(dst) or not osp.splitext(dst)[1]:
89
+ folder = dst
90
+ filename = src.rpartition("/")[2].split("?")[0]
91
+ else:
92
+ folder = osp.dirname(dst)
93
+ filename = osp.basename(dst)
94
+
95
+ local_path = download_url(src, folder, filename=filename, log=log)
96
+ if extract:
97
+ extract_archive(local_path, folder)
98
+ else:
99
+ if osp.isdir(src):
100
+ shutil.copytree(src, dst)
101
+ else:
102
+ os.makedirs(osp.dirname(dst) if osp.splitext(dst)[1] else dst, exist_ok=True)
103
+ shutil.copy(src, dst)
104
+ local_path = osp.join(dst, osp.basename(src)) if osp.isdir(dst) else dst
105
+ if extract:
106
+ extract_archive(local_path, osp.dirname(local_path))
107
+
108
+
109
+ def rm(path: str):
110
+ if osp.isdir(path):
111
+ shutil.rmtree(path)
112
+ elif osp.exists(path):
113
+ os.remove(path)
114
+
115
+
116
+ def glob_files(pattern: str) -> List[str]:
117
+ return sorted(glob.glob(pattern))