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,599 @@
1
+ """Hugging Face Hub integration mixin and methods for K3-Node models and tasks."""
2
+
3
+ import json
4
+ import os
5
+ from pathlib import Path
6
+ from typing import Any, Dict, Optional, Type, TypeVar, Union
7
+ import inspect
8
+ import numpy as np
9
+
10
+ from k3_node.hub.model_card import generate_model_card
11
+
12
+ T = TypeVar("T", bound="K3NodeHubMixin")
13
+
14
+
15
+ class K3NodeHubMixin:
16
+ r"""Mixin class providing seamless saving, loading, and publishing of
17
+ K3-Node GNN models and tasks to/from the Hugging Face Hub.
18
+
19
+ Methods:
20
+ save_pretrained: Saves model weights, config, and Model Card to a local directory.
21
+ from_pretrained: Loads a model from a local folder or Hugging Face Hub repository.
22
+ push_to_hub: Automatically saves and pushes the model to a Hugging Face Hub repo.
23
+ predict: Runs inference on graph data (PyG Data, molecular structures, dicts, or tensors).
24
+ """
25
+
26
+ def _get_config(self) -> Dict[str, Any]:
27
+ r"""Extracts serializable architecture and task hyperparameters."""
28
+ config: Dict[str, Any] = {
29
+ "model_class": self.__class__.__name__,
30
+ }
31
+ # Include task_type for tasks
32
+ if (
33
+ hasattr(self, "backbone")
34
+ or self.__class__.__name__.endswith("Classifier")
35
+ or self.__class__.__name__.endswith("Regressor")
36
+ or self.__class__.__name__.endswith("Predictor")
37
+ ):
38
+ config["task_type"] = self.__class__.__name__
39
+
40
+ for attr in [
41
+ "backbone",
42
+ "in_channels",
43
+ "hidden_channels",
44
+ "out_channels",
45
+ "num_classes",
46
+ "num_layers",
47
+ "num_filters",
48
+ "num_interactions",
49
+ "num_gaussians",
50
+ "cutoff",
51
+ "max_num_neighbors",
52
+ "readout",
53
+ "dipole",
54
+ "mean",
55
+ "std",
56
+ "units",
57
+ "nblocks",
58
+ "dim_atom_embedding",
59
+ "dim_bond_embedding",
60
+ "dim_angle_embedding",
61
+ "num_blocks",
62
+ "pooling",
63
+ "decoder",
64
+ "loss_name",
65
+ "act",
66
+ "multi_label",
67
+ ]:
68
+ if hasattr(self, attr):
69
+ val = getattr(self, attr)
70
+ if val is not None:
71
+ if isinstance(val, (int, float, str, bool, list, dict)):
72
+ config[attr] = val
73
+ elif isinstance(val, tuple):
74
+ config[attr] = list(val)
75
+ elif hasattr(val, "__name__"):
76
+ config[attr] = val.__name__
77
+
78
+ # Handle dropout specifically
79
+ if hasattr(self, "dropout_p") and isinstance(self.dropout_p, (int, float)):
80
+ config["dropout"] = float(self.dropout_p)
81
+ elif hasattr(self, "dropout") and isinstance(self.dropout, (int, float)):
82
+ config["dropout"] = float(self.dropout)
83
+
84
+ # Signature inspection for any remaining constructor arguments
85
+ try:
86
+ sig = inspect.signature(self.__class__.__init__)
87
+ for param_name, param in sig.parameters.items():
88
+ if param_name in ("self", "args", "kwargs", "name"):
89
+ continue
90
+ if param_name not in config and hasattr(self, param_name):
91
+ val = getattr(self, param_name)
92
+ if isinstance(val, (int, float, str, bool, list, dict)):
93
+ config[param_name] = val
94
+ elif isinstance(val, tuple):
95
+ config[param_name] = list(val)
96
+ elif hasattr(val, "__name__"):
97
+ config[param_name] = val.__name__
98
+ except Exception:
99
+ pass
100
+
101
+ if hasattr(self, "backbone_kwargs") and isinstance(self.backbone_kwargs, dict):
102
+ config["backbone_kwargs"] = self.backbone_kwargs
103
+
104
+ return config
105
+
106
+ def save_pretrained(
107
+ self,
108
+ save_directory: Union[str, Path],
109
+ config: Optional[Dict[str, Any]] = None,
110
+ metrics: Optional[Dict[str, float]] = None,
111
+ dataset_name: Optional[str] = None,
112
+ repo_id: Optional[str] = None,
113
+ license: str = "mit",
114
+ **kwargs,
115
+ ) -> Path:
116
+ r"""Saves model weights, config.json, and README.md (Model Card) to disk.
117
+
118
+ Args:
119
+ save_directory: Directory path to save model files in.
120
+ config: Optional custom configuration dictionary.
121
+ metrics: Optional evaluation metrics dictionary to include in Model Card.
122
+ dataset_name: Optional dataset name for the Model Card.
123
+ repo_id: Optional Hugging Face repository ID.
124
+ license: License identifier. (default: ``"mit"``)
125
+
126
+ Returns:
127
+ Path object of the saved directory.
128
+ """
129
+ save_dir = Path(save_directory)
130
+ save_dir.mkdir(parents=True, exist_ok=True)
131
+
132
+ # 1. Config
133
+ final_config = config or self._get_config()
134
+ config_path = save_dir / "config.json"
135
+ with open(config_path, "w", encoding="utf-8") as f:
136
+ json.dump(final_config, f, indent=2)
137
+
138
+ # 2. Weights
139
+ weights_path = save_dir / "model.weights.h5"
140
+ if getattr(self, "model", None) is not None and hasattr(self.model, "save_weights"):
141
+ self.model.save_weights(str(weights_path))
142
+ elif hasattr(self, "save_weights"):
143
+ self.save_weights(str(weights_path))
144
+ else:
145
+ raise RuntimeError(f"Cannot save weights for model of type '{self.__class__.__name__}'.")
146
+
147
+ # 3. Model Card (README.md)
148
+ task_type = final_config.get("task_type", self.__class__.__name__)
149
+ backbone_str = str(final_config.get("backbone", final_config.get("model_class", "gnn")))
150
+ card_content = generate_model_card(
151
+ task_type=task_type,
152
+ backbone=backbone_str,
153
+ config=final_config,
154
+ metrics=metrics,
155
+ dataset_name=dataset_name,
156
+ repo_id=repo_id,
157
+ license=license,
158
+ )
159
+ readme_path = save_dir / "README.md"
160
+ with open(readme_path, "w", encoding="utf-8") as f:
161
+ f.write(card_content)
162
+
163
+ return save_dir
164
+
165
+ def predict(self, data: Any = None, *args: Any, **kwargs: Any) -> Any:
166
+ r"""Infers predictions on graph or molecular data.
167
+
168
+ Supports PyG / K3-Node ``Data`` objects (extracting ``(z, pos, batch)``
169
+ for molecular models or ``(x, edge_index, ...)`` for standard GNNs),
170
+ dictionaries, tuples of tensors, or direct positional tensors.
171
+
172
+ Args:
173
+ data: Input graph or molecule Data, dict, or tensor.
174
+ *args: Additional positional arguments.
175
+ **kwargs: Additional keyword arguments.
176
+
177
+ Returns:
178
+ Model prediction tensor or array.
179
+ """
180
+ # If this is a BaseTask wrapping a model with its own task predict logic:
181
+ if hasattr(self, "_task_predict"):
182
+ return self._task_predict(data, *args, **kwargs)
183
+
184
+ # 1. Molecular data (z, pos, batch)
185
+ if hasattr(data, "z") and hasattr(data, "pos"):
186
+ batch = getattr(data, "batch", None)
187
+ return self(data.z, data.pos, batch=batch, training=False, **kwargs)
188
+
189
+ # 2. Graph data with node features & edges (x, edge_index, ...)
190
+ if hasattr(data, "x") and hasattr(data, "edge_index"):
191
+ edge_weight = getattr(data, "edge_weight", None)
192
+ edge_attr = getattr(data, "edge_attr", None)
193
+ batch = getattr(data, "batch", None)
194
+ call_kwargs = {}
195
+ try:
196
+ sig = inspect.signature(self.call if hasattr(self, "call") else self.__call__)
197
+ params = sig.parameters
198
+ if "edge_weight" in params and edge_weight is not None:
199
+ call_kwargs["edge_weight"] = edge_weight
200
+ elif "edge_attr" in params and edge_attr is not None:
201
+ call_kwargs["edge_attr"] = edge_attr
202
+ if "batch" in params and batch is not None:
203
+ call_kwargs["batch"] = batch
204
+ except Exception:
205
+ pass
206
+ call_kwargs.update(kwargs)
207
+ return self(data.x, data.edge_index, training=False, **call_kwargs)
208
+
209
+ # 3. Dictionary input (e.g., Materials models CHGNet, MEGNet, M3GNet)
210
+ if isinstance(data, dict):
211
+ if "z" in data and "pos" in data:
212
+ return self(data["z"], data["pos"], batch=data.get("batch"), training=False, **kwargs)
213
+ elif "x" in data and "edge_index" in data:
214
+ return self(data["x"], data["edge_index"], training=False, **kwargs)
215
+ else:
216
+ return self(data, training=False, **kwargs)
217
+
218
+ # 4. Tuple or list of inputs
219
+ if isinstance(data, (tuple, list)):
220
+ return self(*data, training=False, **kwargs)
221
+
222
+ # 5. Direct arguments
223
+ if data is not None and len(args) > 0:
224
+ return self(data, *args, training=False, **kwargs)
225
+ elif data is not None:
226
+ return self(data, training=False, **kwargs)
227
+ else:
228
+ return self(*args, training=False, **kwargs)
229
+
230
+ @classmethod
231
+ def from_pretrained(
232
+ cls: Type[T],
233
+ repo_id_or_path: Union[str, Path],
234
+ revision: Optional[str] = None,
235
+ token: Optional[Union[str, bool]] = None,
236
+ cache_dir: Optional[Union[str, Path]] = None,
237
+ **model_kwargs,
238
+ ) -> T:
239
+ r"""Loads a pretrained K3-Node task or model from a local folder or Hugging Face Hub.
240
+
241
+ Args:
242
+ repo_id_or_path: Local directory path or Hugging Face repo ID (e.g. ``"k3-node/schnet-qm9"``).
243
+ revision: Specific git revision/branch on Hugging Face Hub.
244
+ token: Hugging Face authentication token.
245
+ cache_dir: Cache directory for downloaded Hub files.
246
+ **model_kwargs: Overrides for configuration parameters.
247
+
248
+ Returns:
249
+ Restored and initialized model or task instance with loaded weights.
250
+ """
251
+ repo_path = Path(repo_id_or_path)
252
+
253
+ if repo_path.is_dir():
254
+ config_path = repo_path / "config.json"
255
+ weights_path = repo_path / "model.weights.h5"
256
+ if not config_path.exists():
257
+ raise FileNotFoundError(f"config.json not found in local directory '{repo_path}'.")
258
+ if not weights_path.exists():
259
+ raise FileNotFoundError(f"model.weights.h5 not found in local directory '{repo_path}'.")
260
+ else:
261
+ try:
262
+ from huggingface_hub import hf_hub_download
263
+ except ImportError:
264
+ raise ImportError(
265
+ "The `huggingface_hub` package is required to load models from Hugging Face Hub. "
266
+ "Install it via `pip install huggingface_hub`."
267
+ )
268
+
269
+ repo_id = str(repo_id_or_path)
270
+ config_file = hf_hub_download(
271
+ repo_id=repo_id,
272
+ filename="config.json",
273
+ revision=revision,
274
+ token=token,
275
+ cache_dir=cache_dir,
276
+ )
277
+ weights_file = hf_hub_download(
278
+ repo_id=repo_id,
279
+ filename="model.weights.h5",
280
+ revision=revision,
281
+ token=token,
282
+ cache_dir=cache_dir,
283
+ )
284
+ config_path = Path(config_file)
285
+ weights_path = Path(weights_file)
286
+
287
+ with open(config_path, "r", encoding="utf-8") as f:
288
+ config = json.load(f)
289
+
290
+ # Merge any user overrides
291
+ config.update(model_kwargs)
292
+
293
+ # Resolve target class
294
+ target_cls = cls
295
+ if target_cls.__name__ in ("BaseTask", "K3NodeHubMixin"):
296
+ if "task_type" in config:
297
+ from k3_node import tasks
298
+ target_cls = getattr(tasks, config["task_type"], None)
299
+ if target_cls is None or target_cls.__name__ in ("BaseTask", "K3NodeHubMixin"):
300
+ model_class = config.get("model_class")
301
+ if model_class:
302
+ from k3_node import models
303
+ target_cls = getattr(models, model_class, cls)
304
+
305
+ # Extract arguments compatible with target_cls.__init__
306
+ sig = inspect.signature(target_cls.__init__)
307
+ # Exclude *args/**kwargs: for tasks, ``**backbone_kwargs`` is a var-keyword parameter, and
308
+ # passing ``backbone_kwargs={...}`` would nest the options instead of forwarding them.
309
+ accepted_params = {
310
+ name
311
+ for name, p in sig.parameters.items()
312
+ if name != "self" and p.kind not in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD)
313
+ }
314
+
315
+ init_kwargs = {}
316
+ for k, v in config.items():
317
+ if k in accepted_params and v is not None:
318
+ init_kwargs[k] = v
319
+
320
+ if "backbone_kwargs" in config and "backbone_kwargs" not in accepted_params:
321
+ init_kwargs.update(config["backbone_kwargs"])
322
+
323
+ # Instantiate model or task
324
+ instance = target_cls(**init_kwargs)
325
+
326
+ # Initialize model topology if task
327
+ if hasattr(instance, "_init_model"):
328
+ instance._init_model(None)
329
+
330
+ # Build variables so weights can be loaded
331
+ _build_model_if_needed(instance, config)
332
+
333
+ # Load weights
334
+ if hasattr(instance, "model") and instance.model is not None and hasattr(instance.model, "load_weights"):
335
+ instance.model.load_weights(str(weights_path))
336
+ instance._is_compiled = True
337
+ elif hasattr(instance, "load_weights"):
338
+ instance.load_weights(str(weights_path))
339
+ else:
340
+ raise RuntimeError(f"Instance '{instance}' does not support load_weights.")
341
+
342
+ return instance
343
+
344
+ def push_to_hub(
345
+ self,
346
+ repo_id: str,
347
+ token: Optional[Union[str, bool]] = None,
348
+ private: bool = False,
349
+ commit_message: Optional[str] = None,
350
+ metrics: Optional[Dict[str, float]] = None,
351
+ dataset_name: Optional[str] = None,
352
+ license: str = "mit",
353
+ **kwargs,
354
+ ) -> str:
355
+ r"""Saves the model and pushes it directly to the Hugging Face Hub.
356
+
357
+ Args:
358
+ repo_id: Hugging Face repo ID in format ``"username/model_name"`` or ``"org/model_name"``.
359
+ token: Optional Hugging Face auth token. If not passed, uses cached credentials.
360
+ private: Whether the repository should be private. (default: ``False``)
361
+ commit_message: Optional commit message for the upload.
362
+ metrics: Optional dictionary of evaluation metrics to document in the Model Card.
363
+ dataset_name: Optional dataset name for the Model Card.
364
+ license: License tag. (default: ``"mit"``)
365
+
366
+ Returns:
367
+ Web URL of the repository on Hugging Face Hub.
368
+ """
369
+ try:
370
+ from huggingface_hub import HfApi
371
+ except ImportError:
372
+ raise ImportError(
373
+ "The `huggingface_hub` package is required to push models to Hugging Face Hub. "
374
+ "Install it via `pip install huggingface_hub`."
375
+ )
376
+
377
+ import tempfile
378
+
379
+ api = HfApi(token=token)
380
+ repo_url = api.create_repo(
381
+ repo_id=repo_id,
382
+ repo_type="model",
383
+ private=private,
384
+ exist_ok=True,
385
+ )
386
+
387
+ with tempfile.TemporaryDirectory() as tmpdir:
388
+ self.save_pretrained(
389
+ tmpdir,
390
+ metrics=metrics,
391
+ dataset_name=dataset_name,
392
+ repo_id=repo_id,
393
+ license=license,
394
+ **kwargs,
395
+ )
396
+ api.upload_folder(
397
+ folder_path=tmpdir,
398
+ repo_id=repo_id,
399
+ repo_type="model",
400
+ commit_message=commit_message or "Upload K3-Node GNN model with weights and model card",
401
+ )
402
+
403
+ return f"https://huggingface.co/{repo_id}"
404
+
405
+ def export_onnx(
406
+ self,
407
+ output_path: Union[str, Path],
408
+ dummy_inputs: Optional[Any] = None,
409
+ opset: int = 17,
410
+ dynamic_axes: bool = True,
411
+ **kwargs,
412
+ ) -> Path:
413
+ r"""Exports this model or task to high-performance ONNX format.
414
+
415
+ Args:
416
+ output_path: Target path for the `.onnx` file.
417
+ dummy_inputs: Optional sample input data.
418
+ opset: ONNX operator set version. (default: 17)
419
+ dynamic_axes: Whether graph size dimensions are dynamic. (default: True)
420
+
421
+ Returns:
422
+ Path object pointing to the generated `.onnx` file.
423
+ """
424
+ from k3_node.export.onnx_exporter import export_onnx
425
+ return export_onnx(
426
+ self,
427
+ output_path=output_path,
428
+ dummy_inputs=dummy_inputs,
429
+ opset=opset,
430
+ dynamic_axes=dynamic_axes,
431
+ **kwargs,
432
+ )
433
+
434
+ def export_tflite(
435
+ self,
436
+ output_path: Union[str, Path],
437
+ dummy_inputs: Optional[Any] = None,
438
+ quantization: Optional[str] = None,
439
+ **kwargs,
440
+ ) -> Path:
441
+ r"""Exports this model or task to an optimized TensorFlow Lite flatbuffer.
442
+
443
+ Args:
444
+ output_path: Target path for the `.tflite` file.
445
+ dummy_inputs: Optional sample input data.
446
+ quantization: Quantization mode (None, "fp16", "int8_dynamic", "int8_full").
447
+
448
+ Returns:
449
+ Path object pointing to the generated `.tflite` file.
450
+ """
451
+ from k3_node.export.tflite_exporter import export_tflite
452
+ return export_tflite(
453
+ self,
454
+ output_path=output_path,
455
+ dummy_inputs=dummy_inputs,
456
+ quantization=quantization,
457
+ **kwargs,
458
+ )
459
+
460
+ def export_tensorrt(
461
+ self,
462
+ output_path: Union[str, Path],
463
+ dummy_inputs: Optional[Any] = None,
464
+ precision: str = "fp16",
465
+ **kwargs,
466
+ ) -> Path:
467
+ r"""Compiles this model or task into a high-throughput NVIDIA TensorRT engine.
468
+
469
+ Args:
470
+ output_path: Target path for the `.engine` binary.
471
+ dummy_inputs: Optional sample input data.
472
+ precision: Precision mode ("fp32", "fp16", "int8").
473
+
474
+ Returns:
475
+ Path object pointing to the generated TensorRT `.engine` file.
476
+ """
477
+ from k3_node.export.tensorrt_exporter import export_tensorrt
478
+ return export_tensorrt(
479
+ self,
480
+ output_path=output_path,
481
+ dummy_inputs=dummy_inputs,
482
+ precision=precision,
483
+ **kwargs,
484
+ )
485
+
486
+
487
+ def _build_model_if_needed(instance: Any, config: Dict[str, Any]) -> None:
488
+ r"""Builds/initializes weights for tasks and models so weights can be loaded."""
489
+ cls_name = instance.__class__.__name__
490
+
491
+ # Task estimators
492
+ if cls_name == "NodeClassifier":
493
+ in_c = getattr(instance, "in_channels", None) or 16
494
+ dummy_x = np.zeros((2, in_c), dtype="float32")
495
+ dummy_edge = np.zeros((2, 1), dtype="int64")
496
+ instance.model((dummy_x, dummy_edge))
497
+ return
498
+ elif cls_name in ("GraphClassifier", "GraphRegressor"):
499
+ in_c = getattr(instance, "in_channels", None) or 16
500
+ dummy_x = np.zeros((2, in_c), dtype="float32")
501
+ dummy_edge = np.zeros((2, 1), dtype="int64")
502
+ dummy_batch = np.zeros((2,), dtype="int64")
503
+ if hasattr(instance.model, "num_graphs"):
504
+ instance.model.num_graphs = 1
505
+ instance.model((dummy_x, dummy_edge, dummy_batch))
506
+ return
507
+ elif cls_name == "LinkPredictor":
508
+ in_c = getattr(instance, "in_channels", None) or 16
509
+ dummy_x = np.zeros((2, in_c), dtype="float32")
510
+ dummy_edge = np.zeros((2, 1), dtype="int64")
511
+ dummy_label_idx = np.zeros((2, 1), dtype="int64")
512
+ instance.model(((dummy_x, dummy_edge), dummy_label_idx))
513
+ return
514
+ elif hasattr(instance, "model") and instance.model is not None:
515
+ if not getattr(instance.model, "built", False):
516
+ try:
517
+ instance.model.build(None)
518
+ except Exception:
519
+ pass
520
+ return
521
+
522
+ # Direct models
523
+ # Molecular 3D models (SchNet, DimeNet, DimeNetPlusPlus, ViSNet)
524
+ if cls_name in ("SchNet", "DimeNet", "DimeNetPlusPlus", "ViSNet", "GNNFF"):
525
+ try:
526
+ import keras.ops as ops
527
+ z = ops.convert_to_tensor([1, 6], dtype="int32")
528
+ pos = ops.convert_to_tensor([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0]], dtype="float32")
529
+ instance(z, pos)
530
+ return
531
+ except Exception:
532
+ pass
533
+
534
+ # Materials models (CHGNet, MEGNet, M3GNet, TensorNet, SO3Net)
535
+ if cls_name in ("CHGNet", "MEGNet", "M3GNet", "TensorNet", "SO3Net"):
536
+ try:
537
+ crystal = {
538
+ "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),
539
+ "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]], dtype=np.int32),
540
+ "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype=np.int32),
541
+ "node_type": np.array([6, 8, 1, 6], dtype=np.int32),
542
+ "batch": np.array([0, 0, 0, 0], dtype=np.int32),
543
+ "state_attr": np.array([[0.0, 0.0]], dtype=np.float32),
544
+ }
545
+ instance(crystal)
546
+ return
547
+ except Exception:
548
+ pass
549
+
550
+ # Standard GNNs (GCN, GraphSAGE, GIN, GAT, PNA, EdgeCNN, BasicGNN)
551
+ in_c = getattr(instance, "in_channels", None) or config.get("in_channels") or 16
552
+ try:
553
+ import keras.ops as ops
554
+ dummy_x = ops.zeros((2, in_c), dtype="float32")
555
+ dummy_edge = ops.zeros((2, 1), dtype="int64")
556
+ instance(dummy_x, dummy_edge)
557
+ return
558
+ except Exception:
559
+ pass
560
+
561
+ # Generic fallback
562
+ if hasattr(instance, "build"):
563
+ try:
564
+ instance.build(None)
565
+ except Exception:
566
+ pass
567
+
568
+
569
+ # Standalone functional API wrappers
570
+ def save_pretrained(
571
+ model_or_task: Any,
572
+ save_directory: Union[str, Path],
573
+ **kwargs,
574
+ ) -> Path:
575
+ r"""Saves a model or task to disk in Hugging Face Hub format."""
576
+ if hasattr(model_or_task, "save_pretrained"):
577
+ return model_or_task.save_pretrained(save_directory, **kwargs)
578
+ raise TypeError(f"Object of type '{type(model_or_task).__name__}' does not support save_pretrained.")
579
+
580
+
581
+ def from_pretrained(
582
+ repo_id_or_path: Union[str, Path],
583
+ task_cls: Optional[Type[Any]] = None,
584
+ **kwargs,
585
+ ) -> Any:
586
+ r"""Loads a model or task from a local directory or Hugging Face Hub."""
587
+ loader_cls = task_cls or K3NodeHubMixin
588
+ return loader_cls.from_pretrained(repo_id_or_path, **kwargs)
589
+
590
+
591
+ def push_to_hub(
592
+ model_or_task: Any,
593
+ repo_id: str,
594
+ **kwargs,
595
+ ) -> str:
596
+ r"""Pushes a model or task directly to the Hugging Face Hub."""
597
+ if hasattr(model_or_task, "push_to_hub"):
598
+ return model_or_task.push_to_hub(repo_id, **kwargs)
599
+ raise TypeError(f"Object of type '{type(model_or_task).__name__}' does not support push_to_hub.")