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
k3_node/hub/hub_mixin.py
ADDED
|
@@ -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.")
|