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.
- k3_node/__init__.py +122 -0
- k3_node/applications/__init__.py +17 -0
- k3_node/applications/bio/__init__.py +21 -0
- k3_node/applications/chemistry/__init__.py +155 -0
- k3_node/applications/materials/__init__.py +127 -0
- k3_node/applications/materials/basis.py +449 -0
- k3_node/applications/materials/chgnet.py +360 -0
- k3_node/applications/materials/core.py +351 -0
- k3_node/applications/materials/grace.py +246 -0
- k3_node/applications/materials/io.py +230 -0
- k3_node/applications/materials/m3gnet.py +462 -0
- k3_node/applications/materials/megnet.py +395 -0
- k3_node/applications/materials/qet.py +220 -0
- k3_node/applications/materials/readout.py +235 -0
- k3_node/applications/materials/so3net.py +234 -0
- k3_node/applications/materials/tensornet.py +381 -0
- k3_node/applications/materials/test_materials.py +167 -0
- k3_node/applications/materials/wrappers.py +95 -0
- k3_node/data/__init__.py +47 -0
- k3_node/data/batch.py +102 -0
- k3_node/data/collate.py +282 -0
- k3_node/data/data.py +532 -0
- k3_node/data/database.py +154 -0
- k3_node/data/dataset.py +182 -0
- k3_node/data/download.py +49 -0
- k3_node/data/extract.py +45 -0
- k3_node/data/feature_store.py +70 -0
- k3_node/data/graph_store.py +92 -0
- k3_node/data/hetero_data.py +374 -0
- k3_node/data/hypergraph_data.py +59 -0
- k3_node/data/in_memory_dataset.py +177 -0
- k3_node/data/makedirs.py +7 -0
- k3_node/data/on_disk_dataset.py +77 -0
- k3_node/data/separate.py +115 -0
- k3_node/data/storage.py +593 -0
- k3_node/data/temporal.py +154 -0
- k3_node/data/test_batch.py +67 -0
- k3_node/data/test_data.py +68 -0
- k3_node/data/test_dataset_and_stores.py +111 -0
- k3_node/data/test_hetero_data.py +33 -0
- k3_node/data/test_temporal_and_hyper.py +32 -0
- k3_node/data/view.py +43 -0
- k3_node/datasets/__init__.py +88 -0
- k3_node/datasets/actor.py +101 -0
- k3_node/datasets/airports.py +84 -0
- k3_node/datasets/amazon.py +66 -0
- k3_node/datasets/ba2motif_dataset.py +73 -0
- k3_node/datasets/ba_shapes.py +81 -0
- k3_node/datasets/bitcoin_otc.py +77 -0
- k3_node/datasets/citation_full.py +81 -0
- k3_node/datasets/coauthor.py +66 -0
- k3_node/datasets/dblp.py +106 -0
- k3_node/datasets/digits.py +63 -0
- k3_node/datasets/email_eu_core.py +60 -0
- k3_node/datasets/entities.py +158 -0
- k3_node/datasets/explainer_dataset.py +101 -0
- k3_node/datasets/facebook.py +51 -0
- k3_node/datasets/fake.py +256 -0
- k3_node/datasets/freebase.py +90 -0
- k3_node/datasets/geometric_shapes.py +69 -0
- k3_node/datasets/github.py +51 -0
- k3_node/datasets/graph_generator/__init__.py +6 -0
- k3_node/datasets/graph_generator/ba_graph.py +20 -0
- k3_node/datasets/graph_generator/base.py +29 -0
- k3_node/datasets/graph_generator/er_graph.py +21 -0
- k3_node/datasets/icews.py +58 -0
- k3_node/datasets/imdb.py +96 -0
- k3_node/datasets/jodie.py +56 -0
- k3_node/datasets/karate.py +56 -0
- k3_node/datasets/lastfm_asia.py +51 -0
- k3_node/datasets/mesh_correspondence.py +50 -0
- k3_node/datasets/molecule_net.py +148 -0
- k3_node/datasets/motif_generator/__init__.py +7 -0
- k3_node/datasets/motif_generator/base.py +29 -0
- k3_node/datasets/motif_generator/custom.py +17 -0
- k3_node/datasets/motif_generator/cycle.py +25 -0
- k3_node/datasets/motif_generator/house.py +27 -0
- k3_node/datasets/movielens.py +55 -0
- k3_node/datasets/planetoid.py +137 -0
- k3_node/datasets/polblogs.py +63 -0
- k3_node/datasets/ppi.py +189 -0
- k3_node/datasets/qm7.py +65 -0
- k3_node/datasets/qm9.py +132 -0
- k3_node/datasets/reddit.py +121 -0
- k3_node/datasets/sbm_dataset.py +165 -0
- k3_node/datasets/seal.py +74 -0
- k3_node/datasets/shape_scenes.py +92 -0
- k3_node/datasets/test_datasets.py +322 -0
- k3_node/datasets/tu_dataset.py +131 -0
- k3_node/datasets/twitch.py +66 -0
- k3_node/datasets/webkb.py +102 -0
- k3_node/datasets/wikics.py +85 -0
- k3_node/datasets/word_net.py +184 -0
- k3_node/etl/__init__.py +37 -0
- k3_node/etl/encoders.py +248 -0
- k3_node/etl/graph_builders.py +270 -0
- k3_node/etl/relational_to_graph.py +201 -0
- k3_node/etl/table_to_graph.py +244 -0
- k3_node/etl/test_etl.py +318 -0
- k3_node/export/__init__.py +15 -0
- k3_node/export/cross_backend.py +172 -0
- k3_node/export/onnx_exporter.py +190 -0
- k3_node/export/runtime.py +254 -0
- k3_node/export/tensorrt_exporter.py +201 -0
- k3_node/export/test_export.py +337 -0
- k3_node/export/tflite_exporter.py +112 -0
- k3_node/hub/__init__.py +29 -0
- k3_node/hub/dataset_hub.py +242 -0
- k3_node/hub/hub_mixin.py +599 -0
- k3_node/hub/model_card.py +133 -0
- k3_node/hub/test_hub.py +419 -0
- k3_node/io/__init__.py +22 -0
- k3_node/io/fs.py +117 -0
- k3_node/io/npz.py +45 -0
- k3_node/io/off.py +29 -0
- k3_node/io/planetoid.py +98 -0
- k3_node/io/tu.py +137 -0
- k3_node/io/txt_array.py +58 -0
- k3_node/layers/__init__.py +14 -0
- k3_node/layers/aggr/__init__.py +70 -0
- k3_node/layers/aggr/attention.py +77 -0
- k3_node/layers/aggr/base.py +403 -0
- k3_node/layers/aggr/basic.py +412 -0
- k3_node/layers/aggr/deep_sets.py +65 -0
- k3_node/layers/aggr/deepsets.py +29 -0
- k3_node/layers/aggr/equilibrium.py +107 -0
- k3_node/layers/aggr/fused.py +43 -0
- k3_node/layers/aggr/gmt.py +89 -0
- k3_node/layers/aggr/gru.py +58 -0
- k3_node/layers/aggr/lcm.py +143 -0
- k3_node/layers/aggr/lstm.py +58 -0
- k3_node/layers/aggr/mlp.py +75 -0
- k3_node/layers/aggr/multi.py +154 -0
- k3_node/layers/aggr/patch_transformer.py +137 -0
- k3_node/layers/aggr/quantile.py +125 -0
- k3_node/layers/aggr/resolver.py +68 -0
- k3_node/layers/aggr/scaler.py +133 -0
- k3_node/layers/aggr/set2set.py +87 -0
- k3_node/layers/aggr/set_transformer.py +107 -0
- k3_node/layers/aggr/sort.py +68 -0
- k3_node/layers/aggr/test_aggr.py +337 -0
- k3_node/layers/aggr/utils.py +210 -0
- k3_node/layers/aggr/variance_preserving.py +54 -0
- k3_node/layers/attention/__init__.py +5 -0
- k3_node/layers/attention/pair_attention.py +448 -0
- k3_node/layers/attention/performer.py +187 -0
- k3_node/layers/attention/polynormer.py +160 -0
- k3_node/layers/attention/qformer.py +143 -0
- k3_node/layers/attention/sgformer.py +106 -0
- k3_node/layers/attention/test_attention.py +68 -0
- k3_node/layers/attention/test_pair_attention.py +91 -0
- k3_node/layers/conv/__init__.py +149 -0
- k3_node/layers/conv/agnn_conv.py +120 -0
- k3_node/layers/conv/antisymmetric_conv.py +94 -0
- k3_node/layers/conv/appnp.py +105 -0
- k3_node/layers/conv/appnp_conv.py +157 -0
- k3_node/layers/conv/arma_conv.py +231 -0
- k3_node/layers/conv/cg_conv.py +92 -0
- k3_node/layers/conv/cheb_conv.py +137 -0
- k3_node/layers/conv/cluster_gcn_conv.py +102 -0
- k3_node/layers/conv/conv.py +100 -0
- k3_node/layers/conv/crystal_conv.py +140 -0
- k3_node/layers/conv/cugraph.py +84 -0
- k3_node/layers/conv/diffusion_conv.py +144 -0
- k3_node/layers/conv/dir_gnn_conv.py +93 -0
- k3_node/layers/conv/dna_conv.py +192 -0
- k3_node/layers/conv/edge_conv.py +107 -0
- k3_node/layers/conv/eg_conv.py +155 -0
- k3_node/layers/conv/fa_conv.py +107 -0
- k3_node/layers/conv/feast_conv.py +126 -0
- k3_node/layers/conv/film_conv.py +143 -0
- k3_node/layers/conv/gat_conv.py +244 -0
- k3_node/layers/conv/gated_graph_conv.py +136 -0
- k3_node/layers/conv/gatv2_conv.py +205 -0
- k3_node/layers/conv/gcn.py +144 -0
- k3_node/layers/conv/gcn2_conv.py +126 -0
- k3_node/layers/conv/gcn_conv.py +135 -0
- k3_node/layers/conv/gen_conv.py +163 -0
- k3_node/layers/conv/general_conv.py +218 -0
- k3_node/layers/conv/gin_conv.py +218 -0
- k3_node/layers/conv/gmm_conv.py +172 -0
- k3_node/layers/conv/gps_conv.py +153 -0
- k3_node/layers/conv/graph_attention.py +262 -0
- k3_node/layers/conv/graph_conv.py +84 -0
- k3_node/layers/conv/gravnet_conv.py +93 -0
- k3_node/layers/conv/han_conv.py +175 -0
- k3_node/layers/conv/heat_conv.py +131 -0
- k3_node/layers/conv/hetero_conv.py +128 -0
- k3_node/layers/conv/hgt_conv.py +218 -0
- k3_node/layers/conv/hypergraph_conv.py +182 -0
- k3_node/layers/conv/le_conv.py +81 -0
- k3_node/layers/conv/lg_conv.py +58 -0
- k3_node/layers/conv/meshcnn_conv.py +84 -0
- k3_node/layers/conv/message_passing.py +451 -0
- k3_node/layers/conv/mf_conv.py +95 -0
- k3_node/layers/conv/mixhop_conv.py +108 -0
- k3_node/layers/conv/nn_conv.py +110 -0
- k3_node/layers/conv/pan_conv.py +100 -0
- k3_node/layers/conv/pdn_conv.py +109 -0
- k3_node/layers/conv/pna_conv.py +177 -0
- k3_node/layers/conv/point_conv.py +101 -0
- k3_node/layers/conv/point_gnn_conv.py +90 -0
- k3_node/layers/conv/point_transformer_conv.py +132 -0
- k3_node/layers/conv/ppf_conv.py +135 -0
- k3_node/layers/conv/ppnp.py +89 -0
- k3_node/layers/conv/res_gated_graph_conv.py +126 -0
- k3_node/layers/conv/rgat_conv.py +251 -0
- k3_node/layers/conv/rgcn_conv.py +321 -0
- k3_node/layers/conv/sage_conv.py +154 -0
- k3_node/layers/conv/sg_conv.py +96 -0
- k3_node/layers/conv/signed_conv.py +100 -0
- k3_node/layers/conv/simple_conv.py +75 -0
- k3_node/layers/conv/spline_conv.py +182 -0
- k3_node/layers/conv/ssg_conv.py +101 -0
- k3_node/layers/conv/supergat_conv.py +195 -0
- k3_node/layers/conv/tag_conv.py +98 -0
- k3_node/layers/conv/test_backend_consistency.py +164 -0
- k3_node/layers/conv/test_conv.py +176 -0
- k3_node/layers/conv/test_conv_pyg.py +566 -0
- k3_node/layers/conv/transformer_conv.py +168 -0
- k3_node/layers/conv/utils.py +403 -0
- k3_node/layers/conv/wl_conv.py +151 -0
- k3_node/layers/conv/x_conv.py +187 -0
- k3_node/layers/dense/__init__.py +40 -0
- k3_node/layers/dense/dense_gat_conv.py +149 -0
- k3_node/layers/dense/dense_gcn_conv.py +117 -0
- k3_node/layers/dense/dense_gin_conv.py +88 -0
- k3_node/layers/dense/dense_graph_conv.py +95 -0
- k3_node/layers/dense/dense_sage_conv.py +85 -0
- k3_node/layers/dense/diff_pool.py +76 -0
- k3_node/layers/dense/dmon_pool.py +223 -0
- k3_node/layers/dense/linear.py +327 -0
- k3_node/layers/dense/mincut_pool.py +92 -0
- k3_node/layers/dense/test_dense.py +377 -0
- k3_node/layers/functional/__init__.py +13 -0
- k3_node/layers/functional/bro.py +49 -0
- k3_node/layers/functional/edge_dropout.py +55 -0
- k3_node/layers/functional/gini.py +44 -0
- k3_node/layers/functional/test_functional.py +34 -0
- k3_node/layers/kge/__init__.py +17 -0
- k3_node/layers/kge/base.py +255 -0
- k3_node/layers/kge/complex.py +98 -0
- k3_node/layers/kge/distmult.py +79 -0
- k3_node/layers/kge/loader.py +50 -0
- k3_node/layers/kge/rotate.py +103 -0
- k3_node/layers/kge/test_kge.py +76 -0
- k3_node/layers/kge/transe.py +96 -0
- k3_node/layers/norm/__init__.py +23 -0
- k3_node/layers/norm/batch_norm.py +328 -0
- k3_node/layers/norm/diff_group_norm.py +141 -0
- k3_node/layers/norm/graph_norm.py +105 -0
- k3_node/layers/norm/graph_size_norm.py +57 -0
- k3_node/layers/norm/instance_norm.py +163 -0
- k3_node/layers/norm/layer_norm.py +245 -0
- k3_node/layers/norm/mean_subtraction_norm.py +57 -0
- k3_node/layers/norm/msg_norm.py +58 -0
- k3_node/layers/norm/pair_norm.py +94 -0
- k3_node/layers/norm/test_norm.py +275 -0
- k3_node/layers/pool/__init__.py +83 -0
- k3_node/layers/pool/approx_knn.py +101 -0
- k3_node/layers/pool/asap.py +173 -0
- k3_node/layers/pool/avg_pool.py +165 -0
- k3_node/layers/pool/cluster_pool.py +168 -0
- k3_node/layers/pool/connect/__init__.py +10 -0
- k3_node/layers/pool/connect/base.py +103 -0
- k3_node/layers/pool/connect/filter_edges.py +113 -0
- k3_node/layers/pool/consecutive.py +30 -0
- k3_node/layers/pool/decimation.py +48 -0
- k3_node/layers/pool/edge_pool.py +189 -0
- k3_node/layers/pool/glob.py +139 -0
- k3_node/layers/pool/graclus.py +66 -0
- k3_node/layers/pool/knn.py +253 -0
- k3_node/layers/pool/max_pool.py +159 -0
- k3_node/layers/pool/mem_pool.py +145 -0
- k3_node/layers/pool/pan_pool.py +144 -0
- k3_node/layers/pool/point_cloud.py +212 -0
- k3_node/layers/pool/pool.py +119 -0
- k3_node/layers/pool/sag_pool.py +174 -0
- k3_node/layers/pool/select/__init__.py +10 -0
- k3_node/layers/pool/select/base.py +112 -0
- k3_node/layers/pool/select/topk.py +206 -0
- k3_node/layers/pool/test_pool.py +456 -0
- k3_node/layers/pool/topk_pool.py +103 -0
- k3_node/layers/pool/voxel_grid.py +70 -0
- k3_node/layers/unpool/__init__.py +9 -0
- k3_node/layers/unpool/knn_interpolate.py +57 -0
- k3_node/layers/unpool/test_unpool.py +31 -0
- k3_node/loader/__init__.py +62 -0
- k3_node/loader/base.py +69 -0
- k3_node/loader/cache.py +68 -0
- k3_node/loader/cluster.py +127 -0
- k3_node/loader/data_list_loader.py +45 -0
- k3_node/loader/dataloader.py +117 -0
- k3_node/loader/dense_data_loader.py +62 -0
- k3_node/loader/dynamic_batch_sampler.py +93 -0
- k3_node/loader/graph_saint.py +188 -0
- k3_node/loader/hgt_loader.py +90 -0
- k3_node/loader/imbalanced_sampler.py +87 -0
- k3_node/loader/keras_dataset.py +334 -0
- k3_node/loader/link_loader.py +179 -0
- k3_node/loader/link_neighbor_loader.py +202 -0
- k3_node/loader/mixin.py +190 -0
- k3_node/loader/neighbor_loader.py +159 -0
- k3_node/loader/neighbor_sampler.py +167 -0
- k3_node/loader/node_loader.py +185 -0
- k3_node/loader/prefetch.py +115 -0
- k3_node/loader/random_node_loader.py +89 -0
- k3_node/loader/sampler_utils.py +499 -0
- k3_node/loader/shadow.py +115 -0
- k3_node/loader/temporal_dataloader.py +98 -0
- k3_node/loader/test_dataloader.py +113 -0
- k3_node/loader/test_keras_dataset.py +221 -0
- k3_node/loader/test_neighbor_loader.py +122 -0
- k3_node/loader/test_sampler_utils.py +82 -0
- k3_node/loader/test_samplers.py +96 -0
- k3_node/loader/test_subgraph_loaders.py +89 -0
- k3_node/loader/utils.py +232 -0
- k3_node/loader/zip_loader.py +88 -0
- k3_node/metrics.py +94 -0
- k3_node/models/__init__.py +424 -0
- k3_node/models/attentive_fp.py +232 -0
- k3_node/models/attract_repel.py +108 -0
- k3_node/models/autoencoder.py +318 -0
- k3_node/models/basic_gnn.py +443 -0
- k3_node/models/bio/__init__.py +4 -0
- k3_node/models/captum.py +52 -0
- k3_node/models/chemistry/__init__.py +4 -0
- k3_node/models/correct_and_smooth.py +146 -0
- k3_node/models/deep_graph_infomax.py +113 -0
- k3_node/models/deepgcn.py +121 -0
- k3_node/models/dimenet.py +737 -0
- k3_node/models/dimenet_utils.py +153 -0
- k3_node/models/gnnff.py +263 -0
- k3_node/models/gps_model.py +1122 -0
- k3_node/models/gpse.py +638 -0
- k3_node/models/graph_unet.py +199 -0
- k3_node/models/graphmae2.py +954 -0
- k3_node/models/graphormer.py +1258 -0
- k3_node/models/graphormer_3d.py +868 -0
- k3_node/models/grover.py +1066 -0
- k3_node/models/jumping_knowledge.py +200 -0
- k3_node/models/label_prop.py +110 -0
- k3_node/models/lightgcn.py +171 -0
- k3_node/models/linkx.py +181 -0
- k3_node/models/lpformer.py +404 -0
- k3_node/models/mask_label.py +114 -0
- k3_node/models/materials/__init__.py +33 -0
- k3_node/models/meta.py +133 -0
- k3_node/models/metapath2vec.py +234 -0
- k3_node/models/mlp.py +264 -0
- k3_node/models/mole_bert.py +379 -0
- k3_node/models/neural_fingerprint.py +95 -0
- k3_node/models/node2vec.py +213 -0
- k3_node/models/pmlp.py +157 -0
- k3_node/models/polynormer.py +229 -0
- k3_node/models/rect.py +93 -0
- k3_node/models/renet.py +221 -0
- k3_node/models/rev_gnn.py +128 -0
- k3_node/models/schnet.py +484 -0
- k3_node/models/sgformer.py +195 -0
- k3_node/models/signed_gcn.py +185 -0
- k3_node/models/test_attentive_fp.py +32 -0
- k3_node/models/test_attract_repel.py +33 -0
- k3_node/models/test_autoencoder.py +119 -0
- k3_node/models/test_basic_gnn.py +102 -0
- k3_node/models/test_correct_and_smooth.py +40 -0
- k3_node/models/test_deep_graph_infomax.py +68 -0
- k3_node/models/test_deepgcn.py +21 -0
- k3_node/models/test_dimenet.py +86 -0
- k3_node/models/test_domain_apis.py +138 -0
- k3_node/models/test_gnnff.py +24 -0
- k3_node/models/test_gps_model.py +271 -0
- k3_node/models/test_gpse.py +34 -0
- k3_node/models/test_graph_unet.py +26 -0
- k3_node/models/test_graphmae2.py +226 -0
- k3_node/models/test_graphormer.py +233 -0
- k3_node/models/test_graphormer3d.py +163 -0
- k3_node/models/test_grover.py +287 -0
- k3_node/models/test_jumping_knowledge.py +129 -0
- k3_node/models/test_label_prop.py +37 -0
- k3_node/models/test_lightgcn.py +38 -0
- k3_node/models/test_linkx.py +31 -0
- k3_node/models/test_lpformer.py +22 -0
- k3_node/models/test_mask_label.py +90 -0
- k3_node/models/test_meta.py +159 -0
- k3_node/models/test_metapath2vec.py +45 -0
- k3_node/models/test_mlp.py +62 -0
- k3_node/models/test_mole_bert.py +164 -0
- k3_node/models/test_neural_fingerprint.py +13 -0
- k3_node/models/test_node2vec.py +57 -0
- k3_node/models/test_pmlp.py +81 -0
- k3_node/models/test_polynormer.py +104 -0
- k3_node/models/test_rect.py +23 -0
- k3_node/models/test_renet.py +32 -0
- k3_node/models/test_rev_gnn.py +24 -0
- k3_node/models/test_schnet.py +43 -0
- k3_node/models/test_sgformer.py +48 -0
- k3_node/models/test_signed_gcn.py +28 -0
- k3_node/models/test_tgn.py +77 -0
- k3_node/models/test_unimol.py +179 -0
- k3_node/models/test_unimol2.py +114 -0
- k3_node/models/test_unimol_plus.py +131 -0
- k3_node/models/test_visnet.py +44 -0
- k3_node/models/tgn.py +382 -0
- k3_node/models/unimol.py +1156 -0
- k3_node/models/unimol2.py +616 -0
- k3_node/models/unimol_docking_v2.py +301 -0
- k3_node/models/unimol_plus.py +456 -0
- k3_node/models/utils.py +97 -0
- k3_node/models/visnet.py +759 -0
- k3_node/ops/__init__.py +4 -0
- k3_node/ops/conv.py +56 -0
- k3_node/ops/creation.py +43 -0
- k3_node/ops/graph.py +27 -0
- k3_node/ops/host.py +41 -0
- k3_node/ops/matmul.py +49 -0
- k3_node/ops/numpy.py +24 -0
- k3_node/ops/segment.py +54 -0
- k3_node/ops/sparse.py +51 -0
- k3_node/rag/__init__.py +49 -0
- k3_node/rag/encoders.py +312 -0
- k3_node/rag/pipeline.py +192 -0
- k3_node/rag/projector.py +184 -0
- k3_node/rag/subgraph.py +270 -0
- k3_node/rag/test_rag.py +347 -0
- k3_node/rag/verbalizer.py +162 -0
- k3_node/tasks/__init__.py +19 -0
- k3_node/tasks/backbone_resolver.py +125 -0
- k3_node/tasks/base.py +67 -0
- k3_node/tasks/graph_classification.py +270 -0
- k3_node/tasks/graph_regression.py +228 -0
- k3_node/tasks/link_prediction.py +306 -0
- k3_node/tasks/node_classification.py +194 -0
- k3_node/tasks/node_regression.py +138 -0
- k3_node/tasks/test_tasks.py +319 -0
- k3_node/test_docstring_examples.py +106 -0
- k3_node/test_training_forwarding.py +116 -0
- k3_node/training.py +115 -0
- k3_node/transforms/__init__.py +166 -0
- k3_node/transforms/base_transform.py +32 -0
- k3_node/transforms/compose.py +58 -0
- k3_node/transforms/general.py +676 -0
- k3_node/transforms/graph.py +1070 -0
- k3_node/transforms/spatial.py +797 -0
- k3_node/transforms/test_random_link_split.py +45 -0
- k3_node/transforms/test_spatial_transforms.py +65 -0
- k3_node/transforms/test_transforms.py +253 -0
- k3_node/transforms/utils.py +102 -0
- k3_node/utils/__init__.py +5 -0
- k3_node/utils/backend_import.py +12 -0
- k3_node/utils/graph.py +286 -0
- k3_node/utils/keras.py +94 -0
- k3_node/utils/random.py +103 -0
- k3_node/utils/smiles.py +235 -0
- k3_node-1.0.0.dist-info/METADATA +284 -0
- k3_node-1.0.0.dist-info/RECORD +459 -0
- k3_node-1.0.0.dist-info/WHEEL +5 -0
- k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
- 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"
|
k3_node/hub/test_hub.py
ADDED
|
@@ -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))
|